2024-09-12 13:44:06 +03:00
#[ cfg(test) ]
mod tests ;
use std ::sync ::Arc ;
use mxlink ::matrix_sdk ::ruma ::{ OwnedEventId , OwnedUserId } ;
use mxlink ::matrix_sdk ::{
2024-10-01 21:06:54 +03:00
deserialized_responses ::TimelineEvent ,
2024-09-12 13:44:06 +03:00
ruma ::events ::{
2024-10-01 21:06:54 +03:00
relation ::Thread ,
2024-09-12 13:44:06 +03:00
room ::message ::{
MessageType , OriginalSyncRoomMessageEvent , Relation , RoomMessageEventContent ,
} ,
AnyMessageLikeEvent , AnyMessageLikeEventContent , AnyTimelineEvent , MessageLikeEvent ,
} ,
Room ,
} ;
2024-09-22 09:34:10 +03:00
use mxlink ::{ MatrixLink , ThreadGetMessagesParams , ThreadInfo } ;
2024-09-12 13:44:06 +03:00
use super ::{ MatrixMessage , MatrixMessageProcessingParams , MatrixMessageType , RoomEventFetcher } ;
2024-10-01 21:06:54 +03:00
use crate ::entity ::{ InteractionContext , InteractionTrigger , MessagePayload } ;
struct DetailedMessagePayload {
is_mentioning_bot : bool ,
message_payload : MessagePayload ,
}
2024-09-12 13:44:06 +03:00
pub async fn get_matrix_messages_in_thread (
matrix_link : MatrixLink ,
room : & Room ,
thread_id : OwnedEventId ,
) -> Result < Vec < MatrixMessage > , mxlink ::matrix_sdk ::Error > {
2024-09-22 09:34:10 +03:00
let messages_native = matrix_link
. threads ( )
. get_messages ( room , thread_id , ThreadGetMessagesParams ::default ( ) )
. await ? ;
2024-09-12 13:44:06 +03:00
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 )
}
2024-10-01 21:06:54 +03:00
pub async fn get_matrix_messages_in_reply_chain (
event_fetcher : & Arc < RoomEventFetcher > ,
room : & Room ,
event_id : OwnedEventId ,
) -> Result < Vec < MatrixMessage > , mxlink ::matrix_sdk ::Error > {
let messages_native =
get_matrix_messages_in_reply_chain_native ( event_fetcher , room , event_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 )
}
async fn get_matrix_messages_in_reply_chain_native (
event_fetcher : & Arc < RoomEventFetcher > ,
room : & Room ,
event_id : OwnedEventId ,
) -> Result < Vec < AnyMessageLikeEvent > , mxlink ::matrix_sdk ::Error > {
let mut next_event_id = Some ( event_id . clone ( ) ) ;
let mut messages : Vec < AnyMessageLikeEvent > = Vec ::new ( ) ;
let mut handled_event_ids : Vec < OwnedEventId > = Vec ::new ( ) ;
while let Some ( next_event_id_in_loop ) = next_event_id {
let event = event_fetcher
. fetch_event_in_room ( & next_event_id_in_loop , room )
. await
. unwrap ( ) ;
if handled_event_ids . contains ( & next_event_id_in_loop ) {
tracing ::warn! (
" Not following loop-causing event: {} " ,
next_event_id_in_loop
) ;
break ;
}
handled_event_ids . push ( next_event_id_in_loop . clone ( ) ) ;
let event_deserialized = event . event . deserialize ( ) ? ;
let AnyTimelineEvent ::MessageLike ( message_like_event ) = event_deserialized else {
tracing ::warn! (
" Not proceeding past non-MessageLike event: {:?} " ,
event_deserialized
) ;
break ;
} ;
next_event_id = match message_like_event . clone ( ) {
AnyMessageLikeEvent ::RoomEncrypted ( _ ) = > None ,
AnyMessageLikeEvent ::RoomMessage ( room_message ) = > {
if let MessageLikeEvent ::Original ( room_message_original ) = room_message {
match room_message_original . content . relates_to {
Some ( Relation ::Reply { in_reply_to } ) = > Some ( in_reply_to . event_id . clone ( ) ) ,
_ = > None ,
}
} else {
None
}
}
_ = > None ,
} ;
messages . push ( message_like_event ) ;
}
messages . reverse ( ) ;
Ok ( messages )
}
pub async fn process_matrix_messages (
2024-09-12 13:44:06 +03:00
messages : & [ MatrixMessage ] ,
params : & MatrixMessageProcessingParams ,
) -> Vec < MatrixMessage > {
let mut messages_filtered : Vec < MatrixMessage > = Vec ::new ( ) ;
for ( i , message ) in messages . iter ( ) . enumerate ( ) {
2024-10-01 21:06:54 +03:00
if ! is_message_from_allowed_sender (
message ,
& params . bot_user_id ,
params . allowed_users . as_deref ( ) ,
) {
2024-09-12 13:44:06 +03:00
continue ;
}
let mut message = message . clone ( ) ;
2024-10-01 21:06:54 +03:00
if i = = 0 & & ! params . first_message_prefixes_to_strip . is_empty ( ) {
2024-09-12 13:44:06 +03:00
let mut message_text = message . message_text . clone ( ) ;
2024-10-01 21:06:54 +03:00
for prefix in & params . 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 ( ) ;
}
// 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 ( & params . bot_user_id )
{
let mut message_text = message . message_text . clone ( ) ;
for prefix in & params . bot_user_prefixes_to_strip {
2024-09-12 13:44:06 +03:00
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
}
2024-10-01 21:06:54 +03:00
/// Tells if the given message is from an allowed sender.
///
/// If allowed_users is None, all messages are allowed.
/// If allowed_users is Some, only messages from the allowed users (and the `bot_user_id`) are allowed.
2024-09-12 13:44:06 +03:00
fn is_message_from_allowed_sender (
matrix_message : & MatrixMessage ,
2024-10-01 21:06:54 +03:00
bot_user_id : & OwnedUserId ,
allowed_users : Option < & [ regex ::Regex ] > ,
2024-09-12 13:44:06 +03:00
) -> bool {
2024-10-01 21:06:54 +03:00
if matrix_message . sender_id = = * bot_user_id {
2024-09-12 13:44:06 +03:00
return true ;
}
2024-10-01 21:06:54 +03:00
if let Some ( allowed_users ) = allowed_users {
if mxidwc ::match_user_id ( matrix_message . sender_id . as_str ( ) , allowed_users ) {
return true ;
}
} else {
// No allowed users configured, so all messages are allowed
2024-09-12 13:44:06 +03:00
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 ,
} ;
2024-10-01 21:06:54 +03:00
let is_reply = matches! ( room_message . relates_to , Some ( Relation ::Reply { .. } ) ) ;
let text = if is_reply {
// For regular replies, we need to strip the fallback-for-rich replies part.
// See: https://spec.matrix.org/v1.11/client-server-api/#fallbacks-for-rich-replies
strip_rich_reply_fallback_text ( & text )
} else {
text
} ;
let mentioned_users = room_message
. mentions
. map ( | m | m . user_ids . iter ( ) . map ( | u | u . to_owned ( ) ) . collect ( ) )
. unwrap_or ( vec! [ ] ) ;
2024-09-12 13:44:06 +03:00
Some ( MatrixMessage {
2024-10-01 21:06:54 +03:00
sender_id : matrix_native_event . sender ( ) . to_owned ( ) ,
2024-09-12 13:44:06 +03:00
message_type : if is_notice {
MatrixMessageType ::Notice
} else {
MatrixMessageType ::Text
} ,
message_text : text ,
2024-10-01 21:06:54 +03:00
mentioned_users ,
2024-09-12 13:44:06 +03:00
} )
}
2024-10-01 21:06:54 +03:00
/// Determines the interaction context for an incoming (new) room event.
///
/// This context is created based on the "newest message" (`current_event`), which is:
/// - either a top-level message, which may or may not be mentioning the bot
/// - this function will inspect the event and will likely start a new threaded conversation
///
/// - or a thread reply
/// - this function will inspect the thread root event and will return the interaction context
/// - if the bot only reacts to prefixed messsages (or mentions), this function may ignore the given thread reply, unless it mentions the bot (which causes a synthetic "first message" to be produced)
/// - if the thread root event is not found, is redacted, or is of some unsupported MessagePayload type, this function will return `None`
///
/// - or an in-room (non-threaded) reply to a room message, which may or may not be mentioning the bot
/// - replies that do not mention the bot cause this function to return `None`
/// - other replies create a interaction context which points to a "first message" which is synthetic
#[ tracing::instrument(name = " determine_interaction_context_for_room_event " , skip_all, fields(room_id = room.room_id().as_str(), event_id = current_event.event_id.as_str())) ]
pub async fn determine_interaction_context_for_room_event (
2024-09-12 13:44:06 +03:00
bot_user_id : & OwnedUserId ,
room : & Room ,
current_event : & OriginalSyncRoomMessageEvent ,
current_event_payload : & MessagePayload ,
event_fetcher : & Arc < RoomEventFetcher > ,
2024-10-01 21:06:54 +03:00
) -> anyhow ::Result < Option < InteractionContext > > {
let current_event_is_mentioning_bot =
is_event_mentioning_bot ( & current_event . content , bot_user_id ) ;
2024-09-12 13:44:06 +03:00
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 ( ) ,
) ;
2024-10-01 21:06:54 +03:00
return Ok ( Some ( InteractionContext {
thread_info ,
trigger : InteractionTrigger {
is_mentioning_bot : current_event_is_mentioning_bot ,
2024-09-12 13:44:06 +03:00
payload : current_event_payload . clone ( ) ,
} ,
} ) ) ;
} ;
2024-10-01 21:06:54 +03:00
match relation {
Relation ::Thread ( thread ) = > {
determine_interaction_context_for_room_event_related_to_thread (
bot_user_id ,
room ,
current_event ,
event_fetcher ,
current_event_is_mentioning_bot ,
thread ,
)
. await
}
Relation ::Reply { in_reply_to } = > {
determine_interaction_context_for_room_event_related_to_reply (
current_event ,
current_event_is_mentioning_bot ,
in_reply_to . event_id . clone ( ) ,
)
. await
}
2024-09-12 13:44:06 +03:00
2024-10-01 21:06:54 +03:00
// This is a replacement or something else. It's not something we support.
_ = > return Ok ( None ) ,
}
}
async fn determine_interaction_context_for_room_event_related_to_thread (
bot_user_id : & OwnedUserId ,
room : & Room ,
current_event : & OriginalSyncRoomMessageEvent ,
event_fetcher : & Arc < RoomEventFetcher > ,
current_event_is_mentioning_bot : bool ,
thread : & Thread ,
) -> anyhow ::Result < Option < InteractionContext > > {
2024-09-12 13:44:06 +03:00
let thread_info = ThreadInfo ::new ( thread . event_id . clone ( ) , current_event . event_id . clone ( ) ) ;
2024-10-01 21:06:54 +03:00
tracing ::trace! (
? current_event_is_mentioning_bot ,
is_thread_root_only = thread_info . is_thread_root_only ( ) ,
" Dealing with a thread reply " ,
) ;
if current_event_is_mentioning_bot & & ! thread_info . is_thread_root_only ( ) {
// If the current event is a thread reply and is mentioning the bot,
// it's probably someone trying to involve us in the threaded conversation.
// See: https://github.com/etkecc/baibot/issues/15
//
// In such cases, we don't care what the thread root event is like or what the current event is like,
// we want text-generation to be triggered for this whole thread regardless.
return Ok ( Some ( InteractionContext {
thread_info ,
trigger : InteractionTrigger {
is_mentioning_bot : true ,
payload : MessagePayload ::SynthethicChatCompletionTriggerInThread ,
} ,
} ) ) ;
}
2024-09-12 13:44:06 +03:00
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 "
) ;
2024-10-01 21:06:54 +03:00
let thread_start_detailed_message_payload = timeline_event_to_detailed_message_payload (
& thread . event_id ,
thread_start_timeline_event ,
thread_info . clone ( ) ,
bot_user_id ,
) ? ;
let Some ( detailed_message_payload ) = thread_start_detailed_message_payload else {
return Ok ( None ) ;
} ;
Ok ( Some ( InteractionContext {
thread_info ,
trigger : InteractionTrigger {
is_mentioning_bot : detailed_message_payload . is_mentioning_bot ,
payload : detailed_message_payload . message_payload ,
} ,
} ) )
}
async fn determine_interaction_context_for_room_event_related_to_reply (
current_event : & OriginalSyncRoomMessageEvent ,
current_event_is_mentioning_bot : bool ,
reply_to_event_id : OwnedEventId ,
) -> anyhow ::Result < Option < InteractionContext > > {
tracing ::trace! ( ? current_event_is_mentioning_bot , " Dealing with a reply " ) ;
if ! current_event_is_mentioning_bot {
// If the current event is not mentioning the bot, we don't care about it.
tracing ::trace! ( " Ignoring reply event which does not mention the bot " ) ;
return Ok ( None ) ;
}
let thread_info = ThreadInfo ::new ( reply_to_event_id . clone ( ) , current_event . event_id . clone ( ) ) ;
Ok ( Some ( InteractionContext {
thread_info ,
trigger : InteractionTrigger {
is_mentioning_bot : true ,
payload : MessagePayload ::SynthethicChatCompletionTriggerForReply ,
} ,
} ) )
}
fn is_event_mentioning_bot (
event_content : & RoomMessageEventContent ,
bot_user_id : & OwnedUserId ,
) -> bool {
// As a fallback, we used to do string matching (`event_content.body().contains(bot_user_id.as_str())`) here as well.
// However, this is unreliable. In 2024+, clients that do not have proper mentions support should get fixed,
// instead of us having to deal with the possibility of false positives.
//
let Some ( mentions ) = & event_content . mentions else {
return false ;
} ;
mentions
. user_ids
. iter ( )
. any ( | user_id | user_id = = bot_user_id )
}
/// Strips the rich reply fallback text from the given text.
/// See: https://spec.matrix.org/v1.11/client-server-api/#fallbacks-for-rich-replies
///
/// Example:
/// ```rust,ignore
/// let text = "> <@admin:example.com> What's the difference between Matrix and XMPP?\n\nAnswer me";
/// let stripped_text = strip_rich_reply_fallback_text(text);
/// assert_eq!(stripped_text, "Answer me");
/// ```
fn strip_rich_reply_fallback_text ( text : & str ) -> String {
let lines = text . lines ( ) ;
let mut stripped_lines = Vec ::new ( ) ;
let mut encountered_non_prefix = false ;
for line in lines {
if ! encountered_non_prefix & & line . starts_with ( " > " ) {
continue ;
} else {
encountered_non_prefix = true ;
stripped_lines . push ( line ) ;
}
}
stripped_lines . join ( " \n " ) . trim ( ) . to_owned ( )
}
fn timeline_event_to_detailed_message_payload (
timeline_event_id : & OwnedEventId ,
timeline_event : TimelineEvent ,
thread_info : ThreadInfo ,
bot_user_id : & OwnedUserId ,
) -> anyhow ::Result < Option < DetailedMessagePayload > > {
let timeline_event_deserialized = match timeline_event . event . deserialize ( ) {
Ok ( value ) = > value ,
Err ( err ) = > {
return Err ( anyhow ::format_err! (
" Failed to deserialize timeline event {}: {:?} " ,
timeline_event_id ,
err
) ) ;
}
} ;
2024-09-12 13:44:06 +03:00
let AnyTimelineEvent ::MessageLike ( thread_start_message_like_event ) =
2024-10-01 21:06:54 +03:00
timeline_event_deserialized
2024-09-12 13:44:06 +03:00
else {
tracing ::trace! (
2024-10-01 21:06:54 +03:00
" Ignoring non-MessageLike timeline event: {:?} " ,
timeline_event_deserialized
2024-09-12 13:44:06 +03:00
) ;
return Ok ( None ) ;
} ;
2024-10-01 21:06:54 +03:00
let ( is_mentioning_bot , message_payload ) = match thread_start_message_like_event {
AnyMessageLikeEvent ::RoomEncrypted ( room_message ) = > {
tracing ::warn! (
" Could not inspect event {} because it failed to decrypt: {:?} " ,
timeline_event_id . clone ( ) ,
room_message
) ;
// There's no way to know and it doesn't matter anyway.
let is_mentioning_bot = false ;
2024-09-12 13:44:06 +03:00
2024-10-01 21:06:54 +03:00
(
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 event message of unknown type " ,
) ;
return Ok ( None ) ;
} ;
2024-09-12 13:44:06 +03:00
2024-10-01 21:06:54 +03:00
let is_mentioning_bot =
is_event_mentioning_bot ( & room_message_original . content , bot_user_id ) ;
2024-09-12 13:44:06 +03:00
2024-10-01 21:06:54 +03:00
( is_mentioning_bot , room_message_payload )
} else {
tracing ::error! ( " Ignoring event message which appears to be redacted " ) ;
2024-09-12 13:44:06 +03:00
return Ok ( None ) ;
}
2024-10-01 21:06:54 +03:00
}
other = > {
tracing ::trace! ( " Ignoring unknown MessageLike event: {:?} " , other ) ;
return Ok ( None ) ;
}
} ;
2024-09-12 13:44:06 +03:00
2024-10-01 21:06:54 +03:00
Ok ( Some ( DetailedMessagePayload {
is_mentioning_bot ,
message_payload ,
2024-09-12 13:44:06 +03:00
} ) )
}
2024-10-01 21:06:54 +03:00
/// Creates a list of prefixes to strip from the beginning of message texts that mention the bot user.
///
/// Different clients do mentions differently.
/// The body text containing the mention usually contains one of:
/// - the full user ID (includes a @ prefix by default)
/// - the localpart (with a @ prefix)
/// - the localpart (without a @ prefix)
/// - the display name (with a @ prefix)
/// - the display name (without a @ prefix)
///
/// Some add a `: ` suffix after the mention.
///
/// There's no guarantee that the mention is at the start even.
/// It being there is most common and we try to strip it from there
/// as best as we can.
pub fn create_list_of_bot_user_prefixes_to_strip (
2024-09-12 13:44:06 +03:00
bot_user_id : & OwnedUserId ,
2024-10-01 21:06:54 +03:00
bot_display_name : & Option < String > ,
) -> Vec < String > {
let bot_user_id_localpart = bot_user_id . localpart ( ) ;
let mut prefixes_to_strip = vec! [
bot_user_id . as_str ( ) . to_owned ( ) ,
format! ( " @ {} " , bot_user_id_localpart ) ,
bot_user_id_localpart . to_owned ( ) ,
] ;
if let Some ( bot_display_name ) = bot_display_name {
prefixes_to_strip . push ( format! ( " @ {} " , bot_display_name ) ) ;
prefixes_to_strip . push ( bot_display_name . to_owned ( ) ) ;
2024-09-12 13:44:06 +03:00
}
2024-10-01 21:06:54 +03:00
prefixes_to_strip . push ( " : " . to_owned ( ) ) ;
prefixes_to_strip
2024-09-12 13:44:06 +03:00
}