Initial commit
This commit is contained in:
17
src/conversation/llm/entity.rs
Normal file
17
src/conversation/llm/entity.rs
Normal 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>,
|
||||
}
|
||||
10
src/conversation/llm/mod.rs
Normal file
10
src/conversation/llm/mod.rs
Normal 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::*;
|
||||
84
src/conversation/llm/tests.rs
Normal file
84
src/conversation/llm/tests.rs
Normal 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());
|
||||
}
|
||||
287
src/conversation/llm/tokenization.rs
Normal file
287
src/conversation/llm/tokenization.rs
Normal 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
|
||||
);
|
||||
}
|
||||
}
|
||||
51
src/conversation/llm/utils.rs
Normal file
51
src/conversation/llm/utils.rs
Normal 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(),
|
||||
})
|
||||
}
|
||||
38
src/conversation/matrix/entity.rs
Normal file
38
src/conversation/matrix/entity.rs
Normal 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
|
||||
}
|
||||
}
|
||||
11
src/conversation/matrix/mod.rs
Normal file
11
src/conversation/matrix/mod.rs
Normal 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::*;
|
||||
56
src/conversation/matrix/room_display_name_fetcher.rs
Normal file
56
src/conversation/matrix/room_display_name_fetcher.rs
Normal 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
|
||||
}
|
||||
}
|
||||
46
src/conversation/matrix/room_event_fetcher.rs
Normal file
46
src/conversation/matrix/room_event_fetcher.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
284
src/conversation/matrix/utils/mod.rs
Normal file
284
src/conversation/matrix/utils/mod.rs
Normal 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, ¶ms.bot_user_id, ¶ms.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 ¶ms.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) = ¤t_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(¤t_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())
|
||||
}
|
||||
}
|
||||
218
src/conversation/matrix/utils/tests.rs
Normal file
218
src/conversation/matrix/utils/tests.rs
Normal 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,
|
||||
);
|
||||
}
|
||||
}
|
||||
34
src/conversation/matrix_llm_bridge.rs
Normal file
34
src/conversation/matrix_llm_bridge.rs
Normal 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, ¶ms.bot_user_id)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
llm_messages.push(llm_message);
|
||||
}
|
||||
|
||||
Ok(Conversation {
|
||||
messages: llm_messages,
|
||||
})
|
||||
}
|
||||
5
src/conversation/mod.rs
Normal file
5
src/conversation/mod.rs
Normal 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;
|
||||
Reference in New Issue
Block a user