Initial commit

This commit is contained in:
Slavi Pantaleev
2024-09-12 13:44:06 +03:00
commit 946aa9d9e9
220 changed files with 26033 additions and 0 deletions

View File

@@ -0,0 +1,17 @@
#[derive(Debug, Clone, PartialEq)]
pub enum Author {
Prompt,
Assistant,
User,
}
#[derive(Debug, Clone)]
pub struct Message {
pub author: Author,
pub message_text: String,
}
#[derive(Debug)]
pub struct Conversation {
pub messages: Vec<Message>,
}

View File

@@ -0,0 +1,10 @@
mod entity;
mod tokenization;
mod utils;
#[cfg(test)]
mod tests;
pub use entity::*;
pub use tokenization::shorten_messages_list_to_context_size;
pub use utils::*;

View File

@@ -0,0 +1,84 @@
use crate::utils::status::create_error_message_text;
use crate::utils::text_to_speech::create_transcribed_message_text;
use super::*;
#[test]
fn test_messages_by_the_bot_are_identified_correctly() {
let bot_user_id = "@bot:example.com";
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(),
};
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!");
}
#[test]
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user(
) {
let bot_user_id = "@bot:example.com";
let source_message_text = "Hello!";
let message_text = create_transcribed_message_text(source_message_text);
assert_ne!(source_message_text, message_text);
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
};
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);
}
#[test]
fn test_notice_error_messages_by_bot_are_ignored() {
let bot_user_id = "@bot:example.com";
let source_message_text = "Some error happened";
let message_text = create_error_message_text(source_message_text);
assert_ne!(source_message_text, message_text);
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id);
assert!(llm_message.is_none());
}
#[test]
fn test_other_notice_messages_by_the_bot_are_ignored() {
// Also see `test_notice_error_messages_by_bot_are_ignored()`.
// That one passes accidentally, because we ignore all messages by the bot that are notices
// (except for speech-to-text-created transcriptions - see `test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user()`).
// This test is to make sure that we don't accidentally start accepting other notice messages.
let bot_user_id = "@bot:example.com";
let message_text = "Something something";
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(),
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id);
assert!(llm_message.is_none());
}

View File

@@ -0,0 +1,287 @@
use tiktoken_rs::get_bpe_from_tokenizer;
use tiktoken_rs::tokenizer;
use tiktoken_rs::CoreBPE;
use super::{Author, Message};
fn get_bpe_for_model(model: &str) -> CoreBPE {
let tokenizer = tokenizer::get_tokenizer(model)
.or_else(|| tokenizer::get_tokenizer("gpt-4"))
.unwrap();
get_bpe_from_tokenizer(tokenizer).unwrap()
}
pub fn shorten_messages_list_to_context_size(
model: &str,
prompt_message: &Option<Message>,
mut messages: Vec<Message>,
max_response_tokens: u32,
max_context_tokens: u32,
) -> Vec<Message> {
// Loading the tokenization data is an expensive process, so
// se construct the BPE instance once and then use it for all messages.
let bpe = get_bpe_for_model(model);
// We want to retain the prompt in all cases, so we always count it first.
// We also always reserve enough tokens for the maximum response we expect.
let mut current_context_length: u32 = if let Some(prompt_message) = prompt_message {
calculate_token_size_for_message(&bpe, model, prompt_message) + max_response_tokens
} else {
0
};
messages.reverse();
let mut messages_to_keep: Vec<Message> = Vec::new();
for message in messages {
let tokens_for_message = calculate_token_size_for_message(&bpe, model, &message);
if current_context_length + tokens_for_message > max_context_tokens {
break;
}
current_context_length += tokens_for_message;
messages_to_keep.push(message);
}
messages_to_keep.reverse();
messages_to_keep
}
/// Calculate the token size of a message for a given model, with a preloaded CoreBPE object.
/// Related to `calculate_token_size_for_model_message`.
fn calculate_token_size_for_message(bpe: &CoreBPE, model: &str, message: &Message) -> u32 {
let (tokens_per_message, tokens_per_name) = if model.starts_with("gpt-3.5") {
(
4, // every message follows <im_start>{role/name}\n{content}<im_end>\n
-1, // if there's a name, the role is omitted
)
} else {
(3, 1)
};
let role_length = match message.author {
Author::Assistant => bpe.encode_with_special_tokens("assistant").len() as i32,
Author::User => bpe.encode_with_special_tokens("user").len() as i32,
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;
(text_length + role_length + tokens_per_message + tokens_per_name) as u32
}
pub mod test {
#[test]
fn message_size_counting_works() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let message = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
};
let tokens = super::calculate_token_size_for_message(&bpe, model, &message);
assert_eq!(8, tokens);
}
#[test]
fn shortening_works_with_english() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let prompt = super::Message {
author: super::Author::Prompt,
message_text: "You are a bot!".to_owned(),
};
let prompt_length = 10;
assert_eq!(
prompt_length,
super::calculate_token_size_for_message(&bpe, model, &prompt)
);
let mut conversation_messages = Vec::new();
let first = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
};
let first_length = 8;
assert_eq!(
first_length,
super::calculate_token_size_for_message(&bpe, model, &first)
);
conversation_messages.push(first);
let second = super::Message {
author: super::Author::Assistant,
message_text: "Hello!".to_owned(),
};
let second_length = 7;
assert_eq!(
second_length,
super::calculate_token_size_for_message(&bpe, model, &second)
);
conversation_messages.push(second);
let third = super::Message {
author: super::Author::User,
message_text: "This is the 3rd message in this conversation. It shall be preserved."
.to_owned(),
};
let third_length = 21;
assert_eq!(
third_length,
super::calculate_token_size_for_message(&bpe, model, &third)
);
conversation_messages.push(third.clone());
let forth = super::Message {
author: super::Author::Assistant,
message_text: "This is yet another message that shall be preserved.".to_owned(),
};
let forth_length = 15;
assert_eq!(
forth_length,
super::calculate_token_size_for_message(&bpe, model, &forth)
);
conversation_messages.push(forth.clone());
assert_eq!(4, conversation_messages.len());
let new_conversation_messages = super::shorten_messages_list_to_context_size(
model,
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
);
}
#[test]
fn shortening_works_with_japanese() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let prompt = super::Message {
author: super::Author::User,
message_text: "あなたはボットです。".to_owned(),
};
let prompt_length = 14;
assert_eq!(
prompt_length,
super::calculate_token_size_for_message(&bpe, model, &prompt)
);
let mut conversation_messages = Vec::new();
let first = super::Message {
author: super::Author::User,
message_text: "こんにちは!".to_owned(),
};
let first_length = 7;
assert_eq!(
first_length,
super::calculate_token_size_for_message(&bpe, model, &first)
);
conversation_messages.push(first);
let second = super::Message {
author: super::Author::Assistant,
message_text: "こんにちは。今日は元気ですか。".to_owned(),
};
let second_length = 15;
assert_eq!(
second_length,
super::calculate_token_size_for_message(&bpe, model, &second)
);
conversation_messages.push(second);
let third = super::Message {
author: super::Author::User,
message_text: "これは第3のメッセージなので、保存されます。".to_owned(),
};
let third_length = 22;
assert_eq!(
third_length,
super::calculate_token_size_for_message(&bpe, model, &third)
);
conversation_messages.push(third.clone());
let forth = super::Message {
author: super::Author::Assistant,
message_text: "これはもう一つの保存されますメッセージです。".to_owned(),
};
let forth_length = 21;
assert_eq!(
forth_length,
super::calculate_token_size_for_message(&bpe, model, &forth)
);
conversation_messages.push(forth.clone());
assert_eq!(4, conversation_messages.len());
let new_conversation_messages = super::shorten_messages_list_to_context_size(
model,
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
);
}
}

View File

@@ -0,0 +1,51 @@
use super::{Author, Message};
use crate::conversation::matrix::{MatrixMessage, MatrixMessageType};
use crate::utils::text_to_speech as text_to_speech_utils;
pub fn convert_matrix_message_to_llm_message(
matrix_message: &MatrixMessage,
bot_user_id: &str,
) -> Option<Message> {
if matrix_message.sender_id == bot_user_id {
return convert_bot_message(matrix_message);
}
convert_user_message(matrix_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),
MatrixMessageType::Notice => convert_bot_notice_message(&matrix_message.message_text),
}
}
fn convert_bot_text_message(text: &str) -> Option<Message> {
Some(Message {
author: Author::Assistant,
message_text: text.to_owned(),
})
}
fn convert_bot_notice_message(text: &str) -> Option<Message> {
// Notice messages sent by the bot are usually transcriptions of previous messages sent by the user.
// Such transcriptions are prefixed with an emoji and blockquoted.
// If we find a notice that doesn't match this pattern, we skip it.
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.
return Some(Message {
author: Author::User,
message_text: text.to_owned(),
});
}
None
}
fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
Some(Message {
author: Author::User,
message_text: matrix_message.message_text.clone(),
})
}

View File

@@ -0,0 +1,38 @@
use regex::Regex;
#[derive(Clone)]
pub struct MatrixMessage {
pub sender_id: String,
pub message_type: MatrixMessageType,
pub message_text: String,
}
#[derive(Clone)]
pub enum MatrixMessageType {
Text,
Notice,
}
#[derive(Default, Clone)]
pub struct MatrixMessageProcessingParams {
pub(crate) bot_user_id: String,
pub(crate) allowed_users: Vec<Regex>,
// If non-empty, these prefixes will be stripped when processing the message
pub(crate) first_message_stripped_prefixes: Vec<String>,
}
impl MatrixMessageProcessingParams {
pub fn new(bot_user_id: String, allowed_users: Vec<Regex>) -> Self {
Self {
bot_user_id,
allowed_users,
..Default::default()
}
}
pub fn with_first_message_stripped_prefixes(mut self, value: Vec<String>) -> Self {
self.first_message_stripped_prefixes = value;
self
}
}

View File

@@ -0,0 +1,11 @@
mod entity;
mod room_display_name_fetcher;
mod room_event_fetcher;
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 utils::*;

View File

@@ -0,0 +1,56 @@
use mxlink::matrix_sdk::ruma::OwnedRoomId;
use mxlink::matrix_sdk::Room;
use mxlink::MatrixLink;
use quick_cache::sync::Cache;
pub struct RoomDisplayNameFetcher {
matrix_link: MatrixLink,
lru_cache: Option<Cache<OwnedRoomId, Option<String>>>,
}
impl RoomDisplayNameFetcher {
pub fn new(matrix_link: MatrixLink, lru_cache_size: Option<usize>) -> Self {
let lru_cache = lru_cache_size.map(Cache::new);
Self {
matrix_link,
lru_cache,
}
}
#[tracing::instrument(skip_all, fields(room_id = room.room_id().as_str()))]
pub async fn own_display_name_in_room(
&self,
room: &Room,
) -> mxlink::matrix_sdk::Result<Option<String>> {
let Some(lru_cache) = &self.lru_cache else {
return self.get_uncached_value(room).await;
};
let guard = lru_cache.get_value_or_guard_async(room.room_id()).await;
match guard {
Ok(value) => {
tracing::debug!("Returning existing cached display name..");
return Ok(value);
}
Err(guard) => {
let value = self.get_uncached_value(room).await?;
let _ = guard.insert(value.clone());
tracing::debug!("Returning now-cached display name");
return Ok(value);
}
}
}
async fn get_uncached_value(&self, room: &Room) -> mxlink::matrix_sdk::Result<Option<String>> {
self.matrix_link
.rooms()
.own_display_name_in_room(room)
.await
}
}

View File

@@ -0,0 +1,46 @@
use mxlink::matrix_sdk::deserialized_responses::TimelineEvent;
use mxlink::matrix_sdk::ruma::OwnedEventId;
use mxlink::matrix_sdk::Room;
use quick_cache::sync::Cache;
pub struct RoomEventFetcher {
lru_cache: Option<Cache<OwnedEventId, TimelineEvent>>,
}
impl RoomEventFetcher {
pub fn new(lru_cache_size: Option<usize>) -> Self {
let lru_cache = lru_cache_size.map(Cache::new);
Self { lru_cache }
}
#[tracing::instrument(skip(self), fields(room_id = room.room_id().as_str(), event_id = event_id.as_str()))]
pub async fn fetch_event_in_room(
&self,
event_id: &OwnedEventId,
room: &Room,
) -> mxlink::matrix_sdk::Result<TimelineEvent> {
let Some(lru_cache) = &self.lru_cache else {
return room.event(event_id).await;
};
let guard = lru_cache.get_value_or_guard_async(event_id).await;
match guard {
Ok(config) => {
tracing::trace!("Returning existing cached event..");
return Ok(config);
}
Err(guard) => {
let event = room.event(event_id).await?;
let _ = guard.insert(event.clone());
tracing::trace!("Returning now-cached event");
return Ok(event);
}
}
}
}

View File

@@ -0,0 +1,284 @@
#[cfg(test)]
mod tests;
use std::sync::Arc;
use mxlink::matrix_sdk::ruma::{OwnedEventId, OwnedUserId};
use mxlink::matrix_sdk::{
ruma::events::{
room::message::{
MessageType, OriginalSyncRoomMessageEvent, Relation, RoomMessageEventContent,
},
AnyMessageLikeEvent, AnyMessageLikeEventContent, AnyTimelineEvent, MessageLikeEvent,
},
Room,
};
use mxlink::{MatrixLink, ThreadInfo};
use super::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType, RoomEventFetcher};
use crate::entity::{MessagePayload, ThreadContext, ThreadContextFirstMessage};
pub async fn get_matrix_messages_in_thread(
matrix_link: MatrixLink,
room: &Room,
thread_id: OwnedEventId,
) -> Result<Vec<MatrixMessage>, mxlink::matrix_sdk::Error> {
let messages_native = matrix_link.threads().get_messages(room, thread_id).await?;
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;
};
messages.push(message);
}
Ok(messages)
}
pub async fn process_matrix_messages_in_thread(
messages: &[MatrixMessage],
params: &MatrixMessageProcessingParams,
) -> Vec<MatrixMessage> {
let mut messages_filtered: Vec<MatrixMessage> = Vec::new();
for (i, message) in messages.iter().enumerate() {
if !is_message_from_allowed_sender(message, &params.bot_user_id, &params.allowed_users) {
continue;
}
let mut message = message.clone();
if i == 0 && !params.first_message_stripped_prefixes.is_empty() {
let mut message_text = message.message_text.clone();
for prefix in &params.first_message_stripped_prefixes {
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();
}
messages_filtered.push(message);
}
messages_filtered
}
fn is_message_from_allowed_sender(
matrix_message: &MatrixMessage,
bot_user_id: &str,
allowed_users: &[regex::Regex],
) -> bool {
if matrix_message.sender_id == bot_user_id {
return true;
}
if mxidwc::match_user_id(&matrix_message.sender_id, allowed_users) {
return true;
}
false
}
pub fn convert_matrix_native_event_to_matrix_message(
matrix_native_event: &AnyMessageLikeEvent,
) -> Option<MatrixMessage> {
let Some(content) = matrix_native_event.original_content() else {
// Redacted message
return None;
};
let AnyMessageLikeEventContent::RoomMessage(room_message) = content else {
// Some state event, etc.
return 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,
};
Some(MatrixMessage {
sender_id: matrix_native_event.sender().to_string(),
message_type: if is_notice {
MatrixMessageType::Notice
} else {
MatrixMessageType::Text
},
message_text: text,
})
}
/// Determines the thread context (relationship within the thread + first thread message payload) for an incoming (new) room event.
/// This room event is assumed to be the "newest message" in the thread (or a top-level message).
/// If the given event is a regular reply (not a thread reply), this function will return `None`.
/// If the given event is a top-level message, this function will consider this event as the start of the thread.
/// If the given event is a thread reply, this function will inspect the thread root event and will return the thread context.
/// If the thread root event is not found, is redacted, or is of some unsupported MessagePayload type, this function will return `None`.
pub async fn determine_thread_context_for_room_event(
bot_user_id: &OwnedUserId,
room: &Room,
current_event: &OriginalSyncRoomMessageEvent,
current_event_payload: &MessagePayload,
event_fetcher: &Arc<RoomEventFetcher>,
) -> anyhow::Result<Option<ThreadContext>> {
let Some(relation) = &current_event.content.relates_to else {
// This is a top-level message. We consider it the start of the thread.
let thread_info = ThreadInfo::new(
current_event.event_id.clone(),
current_event.event_id.clone(),
);
let is_mentioning_bot = is_event_mentioning_bot(&current_event.content, bot_user_id);
return Ok(Some(ThreadContext {
info: thread_info,
first_message: ThreadContextFirstMessage {
is_mentioning_bot,
payload: current_event_payload.clone(),
},
}));
};
let Relation::Thread(thread) = relation else {
// This is a reply or a replacement, etc. It's not a thread.
// We don't care about this.
return Ok(None);
};
let thread_info = ThreadInfo::new(thread.event_id.clone(), current_event.event_id.clone());
let start_time = std::time::Instant::now();
let thread_start_timeline_event = event_fetcher
.fetch_event_in_room(&thread.event_id, room)
.await;
let thread_start_timeline_event = match thread_start_timeline_event {
Ok(value) => value,
Err(err) => {
return Err(anyhow::format_err!(
"Failed to fetch thread start event {}: {:?}",
thread.event_id,
err
));
}
};
let duration = start_time.elapsed();
tracing::trace!(
thread_id = thread.event_id.as_str(),
duration = ?duration,
"Fetched thread start event"
);
let thread_start_timeline_event_deserialized =
match thread_start_timeline_event.event.deserialize() {
Ok(value) => value,
Err(err) => {
return Err(anyhow::format_err!(
"Failed to deserialize thread start event {}: {:?}",
thread.event_id,
err
));
}
};
let AnyTimelineEvent::MessageLike(thread_start_message_like_event) =
thread_start_timeline_event_deserialized
else {
tracing::trace!(
"Ignoring non-MessageLike thread start event: {:?}",
thread_start_timeline_event_deserialized
);
return Ok(None);
};
let (thread_start_message_is_mentioning_bot, thread_start_message_payload) =
match thread_start_message_like_event {
AnyMessageLikeEvent::RoomEncrypted(room_message) => {
tracing::warn!(
"Could not inspect thread start event {} because it failed to decrypt: {:?}",
thread.event_id.clone(),
room_message
);
// There's no way to know and it doesn't matter anyway.
let is_mentioning_bot = false;
(
is_mentioning_bot,
MessagePayload::Encrypted(thread_info.clone()),
)
}
AnyMessageLikeEvent::RoomMessage(room_message) => {
if let MessageLikeEvent::Original(room_message_original) = room_message {
let room_message_payload: Result<MessagePayload, String> =
room_message_original.content.msgtype.clone().try_into();
let Ok(room_message_payload) = room_message_payload else {
tracing::debug!(
msg_type = room_message_original.content.msgtype(),
"Ignoring thread start message of unknown type",
);
return Ok(None);
};
let is_mentioning_bot =
is_event_mentioning_bot(&room_message_original.content, bot_user_id);
(is_mentioning_bot, room_message_payload)
} else {
tracing::error!("Ignoring thread start message which appears to be redacted");
return Ok(None);
}
}
other => {
tracing::trace!(
"Ignoring unknown MessageLike thread start event: {:?}",
other
);
return Ok(None);
}
};
Ok(Some(ThreadContext {
info: thread_info,
first_message: ThreadContextFirstMessage {
is_mentioning_bot: thread_start_message_is_mentioning_bot,
payload: thread_start_message_payload,
},
}))
}
fn is_event_mentioning_bot(
event_content: &RoomMessageEventContent,
bot_user_id: &OwnedUserId,
) -> bool {
if let Some(mentions) = &event_content.mentions {
mentions
.user_ids
.iter()
.any(|user_id| user_id == bot_user_id)
} else {
// For compatibility with clients that do not support the new Mentions specification
// (see https://spec.matrix.org/latest/client-server-api/#user-and-room-mentions),
// we also do string matching here.
//
// It may be even better to match not only against the MXID, but also against the bot's
// room-specific display name.
//
// We may consider dropping this string-matching behavior altogether in the future,
// so improving this compatibility block is not a high priority.
event_content.body().contains(bot_user_id.as_str())
}
}

View File

@@ -0,0 +1,218 @@
use crate::conversation::matrix::{
MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType,
};
#[test]
fn is_message_from_allowed_sender() {
let bot_user_id = "@bot:example.com";
let allowed_user_id = "@user.someone:example.com";
let unallowed_user_id = "@another:example.com";
let bot_message = MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
};
let allowed_user_message = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
};
let unallowed_user_message = MatrixMessage {
sender_id: unallowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
};
let parsed_regex = match mxidwc::parse_pattern("@user.*:example.com") {
Ok(value) => value,
Err(err) => {
panic!("Error parsing regex: {}", err);
}
};
let allowed_users = vec![parsed_regex];
assert!(
super::is_message_from_allowed_sender(&bot_message, bot_user_id, &vec![]),
"Bot message should be allowed"
);
assert!(
super::is_message_from_allowed_sender(&allowed_user_message, bot_user_id, &allowed_users),
"Allowed user message should be allowed"
);
assert!(
!super::is_message_from_allowed_sender(
&unallowed_user_message,
bot_user_id,
&allowed_users
),
"Unallowed user message should be ignored"
);
}
#[tokio::test]
async fn process_matrix_messages_in_thread() {
let bot_user_id = "@bot:example.com";
let allowed_user_id = "@user.someone:example.com";
let unallowed_user_id = "@another:example.com";
let allowed_user_message = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello from the user!".to_owned(),
};
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(),
};
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(),
};
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(),
};
let bot_message = MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello from the bot!".to_owned(),
};
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(),
};
let parsed_regex = match mxidwc::parse_pattern("@user.*:example.com") {
Ok(value) => value,
Err(err) => {
panic!("Error parsing regex: {}", err);
}
};
let allowed_users = vec![parsed_regex];
let message_processing_params_basic =
super::MatrixMessageProcessingParams::new(bot_user_id.to_owned(), allowed_users.clone());
let message_processing_params_with_prefix_stripping =
super::MatrixMessageProcessingParams::new(bot_user_id.to_owned(), allowed_users.clone())
.with_first_message_stripped_prefixes(vec!["!bai".to_owned()]);
struct TestCase {
name: String,
messages: Vec<MatrixMessage>,
message_processing_params: MatrixMessageProcessingParams,
expected_message_texts: Vec<String>,
}
let test_cases = vec![
TestCase {
name: "Messages by unallowed users are ignored".to_owned(),
messages: vec![
allowed_user_message.clone(),
bot_message.clone(),
unallowed_user_message.clone(),
],
message_processing_params: message_processing_params_basic.clone(),
expected_message_texts: vec![
"Hello from the user!".to_owned(),
"Hello from the bot!".to_owned(),
],
},
TestCase {
name: "The first message with a prefix gets stripped if params configure it (regular space)".to_owned(),
messages: vec![
allowed_user_message_with_prefix.clone(),
bot_message.clone(),
allowed_user_message_with_prefix.clone(),
unallowed_user_message.clone(),
],
message_processing_params: message_processing_params_with_prefix_stripping.clone(),
expected_message_texts: vec![
"Hello from the user!".to_owned(),
"Hello from the bot!".to_owned(),
"!bai Hello from the user!".to_owned(),
],
},
TestCase {
name: "The first message with a prefix gets stripped if params configure it (no space)".to_owned(),
messages: vec![
allowed_user_message_with_prefix_no_space.clone(),
bot_message.clone(),
allowed_user_message_with_prefix_no_space.clone(),
unallowed_user_message.clone(),
],
message_processing_params: message_processing_params_with_prefix_stripping.clone(),
expected_message_texts: vec![
"Hello from the user!".to_owned(),
"Hello from the bot!".to_owned(),
"!baiHello from the user!".to_owned(),
],
},
TestCase {
name: "The first message with a prefix gets stripped if params configure it (full-width-space)".to_owned(),
messages: vec![
allowed_user_message_with_prefix_full_width_space.clone(),
bot_message.clone(),
allowed_user_message_with_prefix_full_width_space.clone(),
unallowed_user_message.clone(),
],
message_processing_params: message_processing_params_with_prefix_stripping.clone(),
expected_message_texts: vec![
"Hello from the user!".to_owned(),
"Hello from the bot!".to_owned(),
"!bai Hello from the user!".to_owned(),
],
},
TestCase {
name: "The first message with a prefix remains untouched if params leave it alone"
.to_owned(),
messages: vec![
allowed_user_message_with_prefix.clone(),
bot_message.clone(),
allowed_user_message_with_prefix.clone(),
unallowed_user_message.clone(),
],
message_processing_params: message_processing_params_basic.clone(),
expected_message_texts: vec![
"!bai Hello from the user!".to_owned(),
"Hello from the bot!".to_owned(),
"!bai Hello from the user!".to_owned(),
],
},
];
for test_case in test_cases {
let processed_messages = super::process_matrix_messages_in_thread(
&test_case.messages,
&test_case.message_processing_params,
)
.await;
let processed_message_texts = processed_messages
.iter()
.map(|message| message.message_text.clone())
.collect::<Vec<String>>();
assert_eq!(
processed_message_texts, test_case.expected_message_texts,
"Test case {} failed",
test_case.name,
);
}
}

View File

@@ -0,0 +1,34 @@
use mxlink::matrix_sdk::ruma::OwnedEventId;
use mxlink::MatrixLink;
use super::llm::{convert_matrix_message_to_llm_message, Conversation, Message};
use super::matrix::{
get_matrix_messages_in_thread, process_matrix_messages_in_thread, MatrixMessageProcessingParams,
};
pub async fn create_llm_conversation_for_matrix_thread(
matrix_link: MatrixLink,
room: &mxlink::matrix_sdk::Room,
thread_id: OwnedEventId,
params: &MatrixMessageProcessingParams,
) -> Result<Conversation, mxlink::matrix_sdk::Error> {
let messages = get_matrix_messages_in_thread(matrix_link, room, thread_id).await?;
let messages_filtered = process_matrix_messages_in_thread(&messages, params).await;
let mut llm_messages: Vec<Message> = Vec::new();
for matrix_message in messages_filtered {
let Some(llm_message) =
convert_matrix_message_to_llm_message(&matrix_message, &params.bot_user_id)
else {
continue;
};
llm_messages.push(llm_message);
}
Ok(Conversation {
messages: llm_messages,
})
}

5
src/conversation/mod.rs Normal file
View File

@@ -0,0 +1,5 @@
pub(crate) mod llm;
pub(crate) mod matrix;
mod matrix_llm_bridge;
pub(crate) use matrix_llm_bridge::create_llm_conversation_for_matrix_thread;