Initial commit
This commit is contained in:
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,
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user