diff --git a/README.md b/README.md index 893b651..54698f2 100644 --- a/README.md +++ b/README.md @@ -41,7 +41,7 @@ It's influenced by [chaz](https://github.com/arcuru/chaz), but does **not** use ![Introduction and general usage](./docs/screenshots/introduction-and-general-usage.webp) -You can find more screenshots on the the [🌟 Features](./docs/features.md) and other [πŸ“š Documentation](./docs/README.md) pages, as well as in the [docs/screenshots](./docs/screenshots) directory. +You can find more screenshots on the [🌟 Features](./docs/features.md) and other [πŸ“š Documentation](./docs/README.md) pages, as well as in the [docs/screenshots](./docs/screenshots) directory. ## πŸš€ Getting Started diff --git a/docs/access.md b/docs/access.md index e3fd263..68c11aa 100644 --- a/docs/access.md +++ b/docs/access.md @@ -16,6 +16,7 @@ Users: - βœ… can **invite the bot to rooms** - βœ… can **use all the bot's [features](./features.md)** ([πŸ’¬ Text Generation](./features.md#-text-generation), [🦻 Speech-to-Text](./features.md#-speech-to-text), etc.) by sending room messages +- βœ… can **mention the bot** in threads and reply chains to provoke it to respond to non-user messages (see [πŸ“– Usage / πŸ’¬ Text Generation / On-demand involvement](./usage.md#on-demand-involvement)) - βœ… can **change the bot's configuration in a room** (e.g. `!bai config room ...` commands) - ❌ cannot **change the bot's global configuration** (e.g. `!bai config global ...` commands) - ❌ cannot **create new [πŸ€– Agents](./agents.md)** (neither in rooms, nor globally). See [πŸ’Ό Room-local agent managers](#-room-local-agent-managers) for controlling which users can create agents. diff --git a/docs/configuration/text-generation.md b/docs/configuration/text-generation.md index 5b7b428..16689cf 100644 --- a/docs/configuration/text-generation.md +++ b/docs/configuration/text-generation.md @@ -13,7 +13,7 @@ You may also wish to see: In Direct Message rooms with the bot (1:1 rooms), it most usually makes sense for the bot to respond to **all** of your messages, as shown on this [πŸ–ΌοΈ screenshot](../screenshots/text-generation.webp). -In group rooms (with multiple users), it may be more appropriate for the bot to only respond to messages that are **prefixed** with the command prefix (e.g. `!bai`), so that other chat exchange in the room will not trigger it. Such a setup is shown on this [πŸ–ΌοΈ screenshot](../screenshots/text-generation-prefix-requirement.webp). +In group rooms (with multiple users), it may be more appropriate for the bot to only respond to messages that are **prefixed** with the command prefix (e.g. `!bai`) or which are [mentioning](https://spec.matrix.org/latest/client-server-api/#user-and-room-mentions) the bot (e.g. `@baibot`), so that other chat exchange in the room will not trigger it. Such a setup is shown on the [πŸ–ΌοΈ On-demand involvement in the room](../screenshots/text-generation-prefix-requirement.webp) screenshot. There are exceptions to these rules, and you can configure the bot to respond only to prefixed messages in a 1:1 room, or to respond to all messages even in a multi-user group room. @@ -27,7 +27,10 @@ By default, the bot is **auto-configured (upon joining a new room)** to use the Example: `!bai config room text-generation set-prefix-requirement-type command_prefix` (this can also be set globally, see [πŸ› οΈ Room Settings](./README.md#room-settings)) -Regardless of this configuration, **the bot will also respond to messages which directly [mention](https://spec.matrix.org/latest/client-server-api/#user-and-room-mentions) the bot** (e.g. `@baibot`), even if they are not prefixed. An example of this can be seen on this [πŸ–ΌοΈ screenshot](../screenshots/text-generation-prefix-requirement.webp). +Regardless of this configuration, **the bot will also respond to messages by allowed [πŸ‘₯ Users](../access.md#-users) which directly [mention](https://spec.matrix.org/latest/client-server-api/#user-and-room-mentions) the bot** (e.g. `@baibot`), even if they are not prefixed. An example of this can be seen on these screenshots: + +- [πŸ–ΌοΈ On-demand involvement in a thread](../screenshots/text-generation-on-demand-thread-involvement.webp) +- [πŸ–ΌοΈ On-demand involvement in a reply chain](../screenshots/text-generation-on-demand-reply-involvement.webp) ### πŸͺ„ Auto Usage diff --git a/docs/features.md b/docs/features.md index 657bd4d..aff8288 100644 --- a/docs/features.md +++ b/docs/features.md @@ -28,6 +28,8 @@ Text Generation is the bot's ability to **respond to users' text messages with t In multi-user (group) rooms, to avoid disturbing the normal conversation between people, the bot is auto-configured to only respond to messages starting with the command prefix (`!bai`) or direct mentions via the [πŸ’¬ Text Generation / πŸ—Ÿ Prefix Requirement Type](./configuration/text-generation.md#-prefix-requirement-type) setting. +Normally, the bot only responds to allowed [πŸ‘₯ Users](./access.md#-users). In certain cases, it's useful for an allowed user to provoke the bot to respond even in foreign threads or reply chains. You can learn more about this feature in the [πŸ“– Usage / πŸ’¬ Text Generation / On-demand involvement](./usage.md#on-demand-involvement) section. + A few other features (like [πŸ—£οΈ Text-to-Speech](#️-text-to-speech) and [🦻 Speech-to-Text](#-speech-to-text)) combine well with Text Generation, so you **don't necessarily need to communicate with the bot via text** (with [Seamless voice interaction](#seamless-voice-interaction), you can communicate only with voice). You may also wish to see: diff --git a/docs/screenshots/text-generation-on-demand-reply-involvement.webp b/docs/screenshots/text-generation-on-demand-reply-involvement.webp new file mode 100644 index 0000000..2905a1d Binary files /dev/null and b/docs/screenshots/text-generation-on-demand-reply-involvement.webp differ diff --git a/docs/screenshots/text-generation-on-demand-thread-involvement.webp b/docs/screenshots/text-generation-on-demand-thread-involvement.webp new file mode 100644 index 0000000..df6bfb7 Binary files /dev/null and b/docs/screenshots/text-generation-on-demand-thread-involvement.webp differ diff --git a/docs/usage.md b/docs/usage.md index 6e74feb..484365f 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -11,10 +11,11 @@ This is related to the [πŸ’¬ Text Generation](./features.md#-text-generation) fe If there's a text-generation handler agent configured, the bot **may** respond to messages sent in the room. -πŸ–ΌοΈ See screenshots of: +See screenshots of: -- the [default Text Generation flow](./screenshots/text-generation.webp) for 1:1 rooms -- the [Text Generation flow in multi-user rooms](./screenshots/text-generation-prefix-requirement.webp) (where the [πŸ—Ÿ Prefix Requirement](./configuration/text-generation.md#-prefix-requirement-type) setting is auto-configured to "required") +- πŸ–ΌοΈ [the default Text Generation flow](./screenshots/text-generation.webp) in 1:1 rooms +- πŸ–ΌοΈ [the Text Generation flow in multi-user rooms](./screenshots/text-generation-prefix-requirement.webp) (where the [πŸ—Ÿ Prefix Requirement](./configuration/text-generation.md#-prefix-requirement-type) setting is auto-configured to "required") +- [on-demand involvement](#on-demand-involvement) Whether the bot responds depends on: @@ -24,12 +25,27 @@ Whether the bot responds depends on: - (🎨 agent capabilities) whether the configured `text-generation` (or `catch-all`) handler agent actually supports text-generation. The provider may lack support for this feature or it may be disabled in the [πŸ€– agents](./agents.md) configuration -- (the [πŸ—Ÿ Prefix Requirement](./configuration/text-generation.md#-prefix-requirement-type) setting) whether a prefix (e.g. `!bai`) is required in front of messages sent to the room. For multi-user rooms, this setting defaults to "required" +- (the [πŸ—Ÿ Prefix Requirement](./configuration/text-generation.md#-prefix-requirement-type) setting) whether a prefix (e.g. `!bai`) or user mention (e.g. `@baibot`) is required for messages sent to the room. For multi-user rooms, this setting defaults to "required". See [on-demand involvement](#on-demand-involvement) for details. Room messages start a threaded conversation where you can continue back-and-forth communication with the bot. Unless you've enabled the [♻️ Context Management](./features.md#️-context-management) feature, all messages will be sent to the agent's API each time. If the context management feature is enabled, older messages may be dropped. +#### On-demand involvement + +In the following 2 cases, it's useful to involve the bot in conversations on-demand: + +1. For multi-user rooms (with the [πŸ—Ÿ Prefix Requirement](./configuration/text-generation.md#-prefix-requirement-type) setting set to "required") +2. In rooms with foreign users (users that are not authorized bot [πŸ‘₯ users](./access.md#-users)) + +In these instances, an allowed [πŸ‘₯ user](./access.md#-users) can also provoke the bot to respond to **any** thread or reply chain by [mentioning](https://spec.matrix.org/latest/client-server-api/#user-and-room-mentions) the bot (e.g. `@baibot Hello!`). The following screenshots demonstrate this behavior: + +- [πŸ–ΌοΈ On-demand involvement in the room](./screenshots/text-generation-prefix-requirement.webp) +- [πŸ–ΌοΈ On-demand involvement in a thread](./screenshots/text-generation-on-demand-thread-involvement.webp) (the Alice user in this example is not an allowed user, yet her messages are still considered as part of the conversation context) +- [πŸ–ΌοΈ On-demand involvement in a reply chain](./screenshots/text-generation-on-demand-reply-involvement.webp) (the Alice user in this example is not an allowed user, yet her messages are still considered as part of the conversation context) + +πŸ’‘ **NOTE**: Normally, the bot **only considers messages from allowed [πŸ‘₯ Users](./access.md#-users)** and ignores all other messages when responding. However, **when the bot is explicitly invoked (via mention)** in a thread or reply chain, **it will consider all messages** in the thread and reply chain (even those from foreign users) as part of the conversation context. + ### πŸ—£οΈ Text-to-Speech diff --git a/src/bot/messaging.rs b/src/bot/messaging.rs index 11967be..fe90ff5 100644 --- a/src/bot/messaging.rs +++ b/src/bot/messaging.rs @@ -11,7 +11,7 @@ use mxlink::{CallbackError, MessageResponseType}; use tracing::Instrument; use crate::{ - conversation::matrix::determine_thread_context_for_room_event, + conversation::matrix::determine_interaction_context_for_room_event, entity::{MessageContext, MessagePayload, RoomConfigContext, TriggerEventInfo}, }; @@ -239,7 +239,7 @@ impl Messaging { } }; - let thread_context = determine_thread_context_for_room_event( + let interaction_context = determine_interaction_context_for_room_event( self.bot.user_id(), &room, &event, @@ -248,16 +248,18 @@ impl Messaging { ) .await; - let thread_context = match thread_context { + let interaction_context = match interaction_context { Ok(value) => value, Err(err) => { - tracing::error!(?err, "Failed to determine thread context for event"); + tracing::error!(?err, "Failed to determine interaction context for event"); return Ok(()); } }; - let Some(thread_context) = thread_context else { - tracing::debug!("Ignoring message with unknown thread context (likely not a threaded message or a top-level message)"); + let Some(interaction_context) = interaction_context else { + tracing::debug!( + "Ignoring message with unknown interaction context (likely not a message for us)" + ); return Ok(()); }; @@ -276,33 +278,13 @@ impl Messaging { room_config_context, self.bot.admin_pattern_regexes().clone(), trigger_event_info, - thread_context.info.clone(), + interaction_context.thread_info.clone(), ); - let bot_display_name = self - .bot - .room_display_name_fetcher() - .own_display_name_in_room(message_context.room()) - .await; - - let bot_display_name = match bot_display_name { - Ok(value) => value, - Err(err) => { - tracing::warn!( - ?err, - "Failed to fetch bot display name. Proceeding without it" - ); - None - } - }; - - // The first event in the thread determines which handler processes the current event. let controller_type = crate::controller::determine_controller( self.bot.command_prefix(), - &thread_context.first_message, + &interaction_context.trigger, &message_context, - self.bot.user_id(), - &bot_display_name, ); tracing::info!(?controller_type, "Determined controller"); @@ -310,7 +292,7 @@ impl Messaging { let _ = room .send_single_receipt( ReceiptType::Read, - thread_context.info.clone().into(), + interaction_context.thread_info.clone().into(), event.event_id.clone(), ) .await; diff --git a/src/controller/chat_completion/mod.rs b/src/controller/chat_completion/mod.rs index 04dae46..6ab7f9e 100644 --- a/src/controller/chat_completion/mod.rs +++ b/src/controller/chat_completion/mod.rs @@ -18,13 +18,28 @@ use crate::entity::roomconfig::{ use crate::entity::MessagePayload; use crate::strings; use crate::utils::text_to_speech::create_transcribed_message_text; -use crate::{conversation::create_llm_conversation_for_matrix_thread, entity::MessageContext, Bot}; +use crate::{ + conversation::{ + create_llm_conversation_for_matrix_reply_chain, create_llm_conversation_for_matrix_thread, + matrix::create_list_of_bot_user_prefixes_to_strip, + }, + entity::MessageContext, + Bot, +}; #[derive(Debug, PartialEq)] pub enum ChatCompletionControllerType { - ViaText { prefixes_to_strip: Vec }, + // Invoked via a command prefix (e.g. `!bai Hello!`) + TextCommand, + // Invoked via a mention (e.g. `@baibot Hello!`) + TextMention, + // Invoked via a direct message (e.g. `Hello!`) + TextDirect, - ViaAudio, + Audio, + + ThreadMention, + ReplyMention, } struct TextToSpeechEligiblePayload { @@ -125,7 +140,15 @@ pub async fn handle( None }; - let response_type = MessageResponseType::InThread(message_context.thread_info().clone()); + let response_type = match controller_type { + // When we're triggered via a reply mention, we reply to the message that triggered us. + ChatCompletionControllerType::ReplyMention => { + MessageResponseType::Reply(message_context.thread_info().last_event_id.clone()) + } + + // In all other cases, we're dealing with a threaded conversation, so we reply in the thread. + _ => MessageResponseType::InThread(message_context.thread_info().clone()), + }; let text_to_speech_eligible_payload = handle_stage_text_generation( bot, @@ -353,24 +376,78 @@ async fn handle_stage_text_generation( ) .await?; - let prefixes_to_strip = match controller_type { - ChatCompletionControllerType::ViaText { prefixes_to_strip } => prefixes_to_strip.clone(), - ChatCompletionControllerType::ViaAudio => vec![], + // We only strip text from the first message if we're invoked via a command prefix. + // Otherwise, we do bot-user mentions stripping on all messages below. + let first_message_prefixes_to_strip = match controller_type { + ChatCompletionControllerType::TextCommand => vec![bot.command_prefix().to_owned()], + _ => vec![], }; - let params = MatrixMessageProcessingParams::new( - bot.user_id().as_str().to_owned(), - message_context.combined_admin_and_user_regexes(), - ) - .with_first_message_stripped_prefixes(prefixes_to_strip); + let bot_display_name = bot + .room_display_name_fetcher() + .own_display_name_in_room(message_context.room()) + .await; - let conversation = create_llm_conversation_for_matrix_thread( - matrix_link.clone(), - message_context.room(), - message_context.thread_info().root_event_id.clone(), - ¶ms, - ) - .await; + let bot_display_name = match bot_display_name { + Ok(value) => value, + Err(err) => { + tracing::warn!( + ?err, + "Failed to fetch bot display name. Proceeding without it" + ); + None + } + }; + + let bot_user_prefixes_to_strip = + create_list_of_bot_user_prefixes_to_strip(bot.user_id(), &bot_display_name); + + let allowed_users = match controller_type { + // Regular chat completion only operates on messages from allowed users. + ChatCompletionControllerType::TextCommand + | ChatCompletionControllerType::TextMention + | ChatCompletionControllerType::TextDirect + | ChatCompletionControllerType::Audio => { + Some(message_context.combined_admin_and_user_regexes()) + } + + // When we're triggered via an explicit mention (thread or reply), we wish to operate against the mention's whole context + // (the whole thread or the whole reply chain upward of the message that triggered us). + // + // This is to allow admins and users to trigger text-generation for other users' messages. + // When we're dragged into a conversation by a known (to us) user, we'd like to process all messages in the conversation, + // not just those from allowed users. + ChatCompletionControllerType::ThreadMention + | ChatCompletionControllerType::ReplyMention => None, + }; + + let params = MatrixMessageProcessingParams::new(bot.user_id().to_owned(), allowed_users) + .with_first_message_prefixes_to_strip(first_message_prefixes_to_strip) + .with_bot_user_prefixes_to_strip(bot_user_prefixes_to_strip); + + let conversation = match controller_type { + // When we're triggered via a reply mention, the context is the whole reply chain upward of the message that triggered us. + ChatCompletionControllerType::ReplyMention => { + create_llm_conversation_for_matrix_reply_chain( + &bot.room_event_fetcher().clone(), + message_context.room(), + message_context.thread_info().last_event_id.clone(), + ¶ms, + ) + .await + } + + // Everything else is happening in a thread, so the context is the whole thread. + _ => { + create_llm_conversation_for_matrix_thread( + matrix_link.clone(), + message_context.room(), + message_context.thread_info().root_event_id.clone(), + ¶ms, + ) + .await + } + }; let conversation = match conversation { Ok(conversation) => conversation, @@ -565,11 +642,15 @@ async fn handle_stage_speech_to_text_actual_transcribing( // // Regardless of how we post this message, it will be posted as a notice, // which can indicate to the bot (for potential future text-generation purposes) that this message is not a bot message. - let (transcribed_text, annotate_message_with_reaction) = if let MessageResponseType::InThread(_) = response_type { - (create_transcribed_message_text(&speech_to_text_result.text), false) - } else { - (speech_to_text_result.text, true) - }; + let (transcribed_text, annotate_message_with_reaction) = + if let MessageResponseType::InThread(_) = response_type { + ( + create_transcribed_message_text(&speech_to_text_result.text), + false, + ) + } else { + (speech_to_text_result.text, true) + }; let result = bot .messaging() diff --git a/src/controller/determination/mod.rs b/src/controller/determination/mod.rs index 737235d..7fdd117 100644 --- a/src/controller/determination/mod.rs +++ b/src/controller/determination/mod.rs @@ -1,13 +1,11 @@ #[cfg(test)] mod tests; -use mxlink::matrix_sdk::ruma::OwnedUserId; - use super::chat_completion::ChatCompletionControllerType; use crate::{ entity::{ - roomconfig::TextGenerationPrefixRequirementType, MessageContext, MessagePayload, - ThreadContextFirstMessage, + roomconfig::TextGenerationPrefixRequirementType, InteractionTrigger, MessageContext, + MessagePayload, }, strings, }; @@ -16,12 +14,16 @@ use super::ControllerType; pub fn determine_controller( command_prefix: &str, - first_thread_message: &ThreadContextFirstMessage, + first_thread_message: &InteractionTrigger, message_context: &MessageContext, - bot_user_id: &OwnedUserId, - bot_display_name: &Option, ) -> ControllerType { match &first_thread_message.payload { + MessagePayload::SynthethicChatCompletionTriggerInThread => { + ControllerType::ChatCompletion(ChatCompletionControllerType::ThreadMention) + } + MessagePayload::SynthethicChatCompletionTriggerForReply => { + ControllerType::ChatCompletion(ChatCompletionControllerType::ReplyMention) + } MessagePayload::Text(text_message_content) => { let prefix_requirement_type = message_context .room_config_context() @@ -32,8 +34,6 @@ pub fn determine_controller( &text_message_content.body, prefix_requirement_type, first_thread_message.is_mentioning_bot, - bot_user_id, - bot_display_name, ) } MessagePayload::Encrypted(thread_info) => { @@ -47,7 +47,7 @@ pub fn determine_controller( } } MessagePayload::Audio(_) => { - ControllerType::ChatCompletion(ChatCompletionControllerType::ViaAudio) + ControllerType::ChatCompletion(ChatCompletionControllerType::Audio) } MessagePayload::Reaction { .. } => { panic!("Handling reaction as first message in thread does not make sense") @@ -60,8 +60,6 @@ fn determine_text_controller( text: &str, room_text_generation_prefix_requirement_type: TextGenerationPrefixRequirementType, is_mentioning_bot: bool, - bot_user_id: &OwnedUserId, - bot_display_name: &Option, ) -> ControllerType { let text = text.trim(); @@ -102,53 +100,26 @@ fn determine_text_controller( // Otherwise, it depends on the prefix requirement for text generation - it may be routed for chat completion or ignored. if is_mentioning_bot { - // 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. - 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()); - } - - prefixes_to_strip.push(":".to_owned()); - - return ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip, - }); + return ControllerType::ChatCompletion(ChatCompletionControllerType::TextMention); } + // Regardless of what the prefix requirement is, if we encounter a command prefix, we'll consider it a chat completion via command prefix invokation. + // This is to correctly indicate to the chat completion controller that a command prefix was used, + // so that it can be stripped from the beginning of the message. + if text.starts_with(command_prefix) { + return ControllerType::ChatCompletion(ChatCompletionControllerType::TextCommand); + } + + // We're dealing with a regular message that does not start with a command prefix. + match room_text_generation_prefix_requirement_type { TextGenerationPrefixRequirementType::CommandPrefix => { - if text.starts_with(command_prefix) { - ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![command_prefix.to_owned()], - }) - } else { - ControllerType::Ignore - } + // A prefix is required, but we've already checked (above) that the message does not start with a command prefix. + // It's to be ignored. + ControllerType::Ignore } TextGenerationPrefixRequirementType::No => { - ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![], - }) + ControllerType::ChatCompletion(ChatCompletionControllerType::TextDirect) } } } diff --git a/src/controller/determination/tests.rs b/src/controller/determination/tests.rs index 5655dd6..a8d15da 100644 --- a/src/controller/determination/tests.rs +++ b/src/controller/determination/tests.rs @@ -4,9 +4,6 @@ fn determine_text_controller() { use super::ControllerType; use crate::controller; - let bot_user_id = mxlink::matrix_sdk::ruma::owned_user_id!("@bot:example.com"); - let bot_display_name = "Bot"; - let command_prefix = "!bai"; struct TestCase { @@ -44,9 +41,7 @@ fn determine_text_controller() { is_mentioning_bot: false, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::No, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextCommand), }, TestCase { name: "Access top-level", @@ -110,9 +105,7 @@ fn determine_text_controller() { is_mentioning_bot: false, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::No, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextDirect), }, TestCase { name: "Regular text is ignored when prefix is required", @@ -128,9 +121,7 @@ fn determine_text_controller() { is_mentioning_bot: false, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::CommandPrefix, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec!["!bai".to_owned()], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextCommand), }, TestCase { name: "Command-prefixed text triggers completion even when prefix is not required", @@ -138,58 +129,35 @@ fn determine_text_controller() { is_mentioning_bot: false, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::No, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextCommand), }, TestCase { - name: "Regular message with bot mention triggers completion stripping bot id and display name (no prefix requirement)", + name: "Regular message with bot mention triggers completion (no prefix requirement)", input: "Regular text goes here", is_mentioning_bot: true, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::No, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![ - "@bot:example.com".to_owned(), - "@bot".to_owned(), - "bot".to_owned(), - "@Bot".to_owned(), - "Bot".to_owned(), - ":".to_owned(), - ], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextMention), }, - // This test case is the same as the one above, just with a different prefix requirement. + // This test case is the same as the one above, just with a different prefix requirement setting. // We expect the same result. TestCase { - name: "Regular message with bot mention triggers completion stripping bot id and display name (command_prefix requirement)", + name: + "Regular message with bot mention triggers completion (command prefix requirement)", input: "Regular text goes here", is_mentioning_bot: true, room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType::CommandPrefix, - expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText { - prefixes_to_strip: vec![ - "@bot:example.com".to_owned(), - "@bot".to_owned(), - "bot".to_owned(), - "@Bot".to_owned(), - "Bot".to_owned(), - ":".to_owned(), - ], - }), + expected: ControllerType::ChatCompletion(ChatCompletionControllerType::TextMention), }, ]; for test_case in test_cases { - let bot_display_name = Some(bot_display_name.to_owned()); - let result = super::determine_text_controller( command_prefix, test_case.input, test_case.room_text_generation_prefix_requirement_type, test_case.is_mentioning_bot, - &bot_user_id, - &bot_display_name, ); assert_eq!(result, test_case.expected, "Test case: {}", test_case.name); } diff --git a/src/controller/image/generation.rs b/src/controller/image/generation.rs index 4cf5843..cdbb9cb 100644 --- a/src/controller/image/generation.rs +++ b/src/controller/image/generation.rs @@ -35,8 +35,8 @@ pub async fn handle_image( }; let params = MatrixMessageProcessingParams::new( - bot.user_id().as_str().to_owned(), - message_context.combined_admin_and_user_regexes(), + bot.user_id().to_owned(), + Some(message_context.combined_admin_and_user_regexes()), ); let conversation = create_llm_conversation_for_matrix_thread( diff --git a/src/conversation/llm/tests.rs b/src/conversation/llm/tests.rs index 2a1d2fc..a50b9cd 100644 --- a/src/conversation/llm/tests.rs +++ b/src/conversation/llm/tests.rs @@ -1,3 +1,5 @@ +use mxlink::matrix_sdk::ruma::OwnedUserId; + use crate::utils::status::create_error_message_text; use crate::utils::text_to_speech::create_transcribed_message_text; @@ -5,15 +7,17 @@ use super::*; #[test] fn test_messages_by_the_bot_are_identified_correctly() { - let bot_user_id = "@bot:example.com"; + let bot_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); 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(), + mentioned_users: vec![], }; - let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id).unwrap(); + 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!"); @@ -22,7 +26,8 @@ fn test_messages_by_the_bot_are_identified_correctly() { #[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 bot_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); let source_message_text = "Hello!"; let message_text = create_transcribed_message_text(source_message_text); @@ -33,9 +38,10 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con sender_id: bot_user_id.to_owned(), message_type: super::super::matrix::MatrixMessageType::Notice, message_text, + mentioned_users: vec![], }; - let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id).unwrap(); + 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); @@ -43,7 +49,8 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con #[test] fn test_notice_error_messages_by_bot_are_ignored() { - let bot_user_id = "@bot:example.com"; + let bot_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); let source_message_text = "Some error happened"; let message_text = create_error_message_text(source_message_text); @@ -54,9 +61,10 @@ fn test_notice_error_messages_by_bot_are_ignored() { sender_id: bot_user_id.to_owned(), message_type: super::super::matrix::MatrixMessageType::Notice, message_text, + mentioned_users: vec![], }; - let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id); + let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id); assert!(llm_message.is_none()); } @@ -68,7 +76,8 @@ fn test_other_notice_messages_by_the_bot_are_ignored() { // (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 bot_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); let message_text = "Something something"; @@ -76,9 +85,10 @@ fn test_other_notice_messages_by_the_bot_are_ignored() { sender_id: bot_user_id.to_owned(), message_type: super::super::matrix::MatrixMessageType::Notice, message_text: message_text.to_owned(), + mentioned_users: vec![], }; - let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id); + let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id); assert!(llm_message.is_none()); } diff --git a/src/conversation/llm/utils.rs b/src/conversation/llm/utils.rs index 89c8296..be6f663 100644 --- a/src/conversation/llm/utils.rs +++ b/src/conversation/llm/utils.rs @@ -1,12 +1,14 @@ +use matrix_sdk::ruma::OwnedUserId; + 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, + bot_user_id: &OwnedUserId, ) -> Option { - if matrix_message.sender_id == bot_user_id { + if matrix_message.sender_id == bot_user_id.as_str() { return convert_bot_message(matrix_message); } diff --git a/src/conversation/matrix/entity.rs b/src/conversation/matrix/entity.rs index bbb41bc..978ba77 100644 --- a/src/conversation/matrix/entity.rs +++ b/src/conversation/matrix/entity.rs @@ -1,10 +1,13 @@ use regex::Regex; +use mxlink::matrix_sdk::ruma::OwnedUserId; + #[derive(Clone)] pub struct MatrixMessage { - pub sender_id: String, + pub sender_id: OwnedUserId, pub message_type: MatrixMessageType, pub message_text: String, + pub mentioned_users: Vec, } #[derive(Clone)] @@ -13,26 +16,42 @@ pub enum MatrixMessageType { Notice, } -#[derive(Default, Clone)] +#[derive(Clone)] pub struct MatrixMessageProcessingParams { - pub(crate) bot_user_id: String, - pub(crate) allowed_users: Vec, + pub(crate) bot_user_id: OwnedUserId, - // If non-empty, these prefixes will be stripped when processing the message - pub(crate) first_message_stripped_prefixes: Vec, + /// The prefixes that will be stripped when processing the messages in the context (thread or reply chain), + /// which are found to be mentioning the bot user (`bot_user_id`). + pub(crate) bot_user_prefixes_to_strip: Vec, + + /// The prefixes that will be stripped when processing the 1st message in the context (thread or reply chain). + pub(crate) first_message_prefixes_to_strip: Vec, + + /// A list of users whose messages are allowed. + /// If None, all messages are allowed. + /// If Some, only messages from the allowed users (and the bot itself, `bot_user_id`) are allowed. + pub(crate) allowed_users: Option>, } impl MatrixMessageProcessingParams { - pub fn new(bot_user_id: String, allowed_users: Vec) -> Self { + pub fn new(bot_user_id: OwnedUserId, allowed_users: Option>) -> Self { Self { bot_user_id, + bot_user_prefixes_to_strip: vec![], + + first_message_prefixes_to_strip: vec![], + allowed_users, - ..Default::default() } } - pub fn with_first_message_stripped_prefixes(mut self, value: Vec) -> Self { - self.first_message_stripped_prefixes = value; + pub fn with_bot_user_prefixes_to_strip(mut self, value: Vec) -> Self { + self.bot_user_prefixes_to_strip = value; + self + } + + pub fn with_first_message_prefixes_to_strip(mut self, value: Vec) -> Self { + self.first_message_prefixes_to_strip = value; self } } diff --git a/src/conversation/matrix/utils/mod.rs b/src/conversation/matrix/utils/mod.rs index 3f28986..cb37ccb 100644 --- a/src/conversation/matrix/utils/mod.rs +++ b/src/conversation/matrix/utils/mod.rs @@ -5,7 +5,9 @@ use std::sync::Arc; use mxlink::matrix_sdk::ruma::{OwnedEventId, OwnedUserId}; use mxlink::matrix_sdk::{ + deserialized_responses::TimelineEvent, ruma::events::{ + relation::Thread, room::message::{ MessageType, OriginalSyncRoomMessageEvent, Relation, RoomMessageEventContent, }, @@ -16,7 +18,12 @@ use mxlink::matrix_sdk::{ use mxlink::{MatrixLink, ThreadGetMessagesParams, ThreadInfo}; use super::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType, RoomEventFetcher}; -use crate::entity::{MessagePayload, ThreadContext, ThreadContextFirstMessage}; +use crate::entity::{InteractionContext, InteractionTrigger, MessagePayload}; + +struct DetailedMessagePayload { + is_mentioning_bot: bool, + message_payload: MessagePayload, +} pub async fn get_matrix_messages_in_thread( matrix_link: MatrixLink, @@ -42,23 +49,123 @@ pub async fn get_matrix_messages_in_thread( Ok(messages) } -pub async fn process_matrix_messages_in_thread( +pub async fn get_matrix_messages_in_reply_chain( + event_fetcher: &Arc, + room: &Room, + event_id: OwnedEventId, +) -> Result, mxlink::matrix_sdk::Error> { + let messages_native = + get_matrix_messages_in_reply_chain_native(event_fetcher, room, event_id).await?; + + let mut messages: Vec = 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, + room: &Room, + event_id: OwnedEventId, +) -> Result, mxlink::matrix_sdk::Error> { + let mut next_event_id = Some(event_id.clone()); + + let mut messages: Vec = Vec::new(); + let mut handled_event_ids: Vec = 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( messages: &[MatrixMessage], params: &MatrixMessageProcessingParams, ) -> Vec { let mut messages_filtered: Vec = Vec::new(); for (i, message) in messages.iter().enumerate() { - if !is_message_from_allowed_sender(message, ¶ms.bot_user_id, ¶ms.allowed_users) { + if !is_message_from_allowed_sender( + message, + ¶ms.bot_user_id, + params.allowed_users.as_deref(), + ) { continue; } let mut message = message.clone(); - if i == 0 && !params.first_message_stripped_prefixes.is_empty() { + if i == 0 && !params.first_message_prefixes_to_strip.is_empty() { let mut message_text = message.message_text.clone(); - for prefix in ¶ms.first_message_stripped_prefixes { + for prefix in ¶ms.first_message_prefixes_to_strip { + if let Some(message_text_stripped) = message_text.strip_prefix(prefix) { + message_text = message_text_stripped.to_owned(); + } + } + + message.message_text = message_text.trim().to_owned(); + } + + // We only strip `bot_user_prefixes_to_strip`-defined prefixes from messages that mention the bot user. + if !params.bot_user_prefixes_to_strip.is_empty() + && message.mentioned_users.contains(¶ms.bot_user_id) + { + let mut message_text = message.message_text.clone(); + + for prefix in ¶ms.bot_user_prefixes_to_strip { if let Some(message_text_stripped) = message_text.strip_prefix(prefix) { message_text = message_text_stripped.to_owned(); } @@ -73,16 +180,25 @@ pub async fn process_matrix_messages_in_thread( messages_filtered } +/// 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. fn is_message_from_allowed_sender( matrix_message: &MatrixMessage, - bot_user_id: &str, - allowed_users: &[regex::Regex], + bot_user_id: &OwnedUserId, + allowed_users: Option<&[regex::Regex]>, ) -> bool { - if matrix_message.sender_id == bot_user_id { + if matrix_message.sender_id == *bot_user_id { return true; } - if mxidwc::match_user_id(&matrix_message.sender_id, allowed_users) { + 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 return true; } @@ -108,30 +224,58 @@ pub fn convert_matrix_native_event_to_matrix_message( _ => return None, }; + 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![]); + Some(MatrixMessage { - sender_id: matrix_native_event.sender().to_string(), + sender_id: matrix_native_event.sender().to_owned(), message_type: if is_notice { MatrixMessageType::Notice } else { MatrixMessageType::Text }, message_text: text, + mentioned_users, }) } -/// 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( +/// 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( bot_user_id: &OwnedUserId, room: &Room, current_event: &OriginalSyncRoomMessageEvent, current_event_payload: &MessagePayload, event_fetcher: &Arc, -) -> anyhow::Result> { +) -> anyhow::Result> { + let current_event_is_mentioning_bot = + is_event_mentioning_bot(¤t_event.content, bot_user_id); + 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( @@ -139,25 +283,73 @@ pub async fn determine_thread_context_for_room_event( 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, + return Ok(Some(InteractionContext { + thread_info, + trigger: InteractionTrigger { + is_mentioning_bot: current_event_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); - }; + 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 + } + // 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, + current_event_is_mentioning_bot: bool, + thread: &Thread, +) -> anyhow::Result> { let thread_info = ThreadInfo::new(thread.event_id.clone(), current_event.event_id.clone()); + 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, + }, + })); + } + let start_time = std::time::Instant::now(); let thread_start_timeline_event = event_fetcher @@ -183,82 +375,46 @@ pub async fn determine_thread_context_for_room_event( "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 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 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 - ); + let Some(detailed_message_payload) = thread_start_detailed_message_payload else { 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 - ); + Ok(Some(InteractionContext { + thread_info, + trigger: InteractionTrigger { + is_mentioning_bot: detailed_message_payload.is_mentioning_bot, + payload: detailed_message_payload.message_payload, + }, + })) +} - // There's no way to know and it doesn't matter anyway. - let is_mentioning_bot = false; +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> { + tracing::trace!(?current_event_is_mentioning_bot, "Dealing with a reply"); - ( - 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 = - room_message_original.content.msgtype.clone().try_into(); + 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 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 thread_info = ThreadInfo::new(reply_to_event_id.clone(), current_event.event_id.clone()); - 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, + Ok(Some(InteractionContext { + thread_info, + trigger: InteractionTrigger { + is_mentioning_bot: true, + payload: MessagePayload::SynthethicChatCompletionTriggerForReply, }, })) } @@ -267,21 +423,157 @@ 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()) - } + // 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> { + 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 + )); + } + }; + + let AnyTimelineEvent::MessageLike(thread_start_message_like_event) = + timeline_event_deserialized + else { + tracing::trace!( + "Ignoring non-MessageLike timeline event: {:?}", + timeline_event_deserialized + ); + return Ok(None); + }; + + 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; + + ( + 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 = + 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); + }; + + 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 event message which appears to be redacted"); + + return Ok(None); + } + } + other => { + tracing::trace!("Ignoring unknown MessageLike event: {:?}", other); + return Ok(None); + } + }; + + Ok(Some(DetailedMessagePayload { + is_mentioning_bot, + message_payload, + })) +} + +/// 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( + bot_user_id: &OwnedUserId, + bot_display_name: &Option, +) -> Vec { + 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()); + } + + prefixes_to_strip.push(":".to_owned()); + + prefixes_to_strip } diff --git a/src/conversation/matrix/utils/tests.rs b/src/conversation/matrix/utils/tests.rs index f7ede76..3383d49 100644 --- a/src/conversation/matrix/utils/tests.rs +++ b/src/conversation/matrix/utils/tests.rs @@ -1,29 +1,35 @@ +use mxlink::matrix_sdk::ruma::OwnedUserId; + 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_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); + let allowed_user_id = OwnedUserId::try_from("@user.someone:example.com").unwrap(); + let unallowed_user_id = OwnedUserId::try_from("@another:example.com").unwrap(); let bot_message = MatrixMessage { sender_id: bot_user_id.to_owned(), message_type: MatrixMessageType::Text, message_text: "Hello!".to_owned(), + mentioned_users: vec![], }; let allowed_user_message = MatrixMessage { sender_id: allowed_user_id.to_owned(), message_type: MatrixMessageType::Text, message_text: "Hello!".to_owned(), + mentioned_users: vec![], }; let unallowed_user_message = MatrixMessage { sender_id: unallowed_user_id.to_owned(), message_type: MatrixMessageType::Text, message_text: "Hello!".to_owned(), + mentioned_users: vec![], }; let parsed_regex = match mxidwc::parse_pattern("@user.*:example.com") { @@ -36,65 +42,96 @@ fn is_message_from_allowed_sender() { let allowed_users = vec![parsed_regex]; assert!( - super::is_message_from_allowed_sender(&bot_message, bot_user_id, &vec![]), + super::is_message_from_allowed_sender(&bot_message, &bot_user_id, Some(&allowed_users)), "Bot message should be allowed" ); assert!( - super::is_message_from_allowed_sender(&allowed_user_message, bot_user_id, &allowed_users), + super::is_message_from_allowed_sender( + &allowed_user_message, + &bot_user_id, + Some(&allowed_users) + ), "Allowed user message should be allowed" ); assert!( !super::is_message_from_allowed_sender( &unallowed_user_message, - bot_user_id, - &allowed_users + &bot_user_id, + Some(&allowed_users), ), "Unallowed user message should be ignored" ); + + assert!( + super::is_message_from_allowed_sender(&unallowed_user_message, &bot_user_id, None,), + "An empty list of allowed users lets everyone through" + ); } #[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"; +async fn process_matrix_messages() { + let bot_user_id = + OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID"); + let allowed_user_id = OwnedUserId::try_from("@user.someone:example.com").unwrap(); + let unallowed_user_id = OwnedUserId::try_from("@another:example.com").unwrap(); let allowed_user_message = MatrixMessage { sender_id: allowed_user_id.to_owned(), message_type: MatrixMessageType::Text, message_text: "Hello from the user!".to_owned(), + mentioned_users: vec![], }; 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(), + mentioned_users: vec![], }; 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(), + mentioned_users: vec![], }; 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(), + mentioned_users: vec![], }; let bot_message = MatrixMessage { sender_id: bot_user_id.to_owned(), message_type: MatrixMessageType::Text, message_text: "Hello from the bot!".to_owned(), + mentioned_users: vec![], + }; + + let allowed_user_message_with_bot_mention = MatrixMessage { + sender_id: allowed_user_id.to_owned(), + message_type: MatrixMessageType::Text, + message_text: "@baibot: Hello from the user!".to_owned(), + mentioned_users: vec![bot_user_id.to_owned()], + }; + + // The message text is the same as above - it mentions the bot, but the actually-mentioned user is another user. + let allowed_user_message_with_another_user_mention = MatrixMessage { + sender_id: allowed_user_id.to_owned(), + message_type: MatrixMessageType::Text, + message_text: allowed_user_message_with_bot_mention.message_text.clone(), + mentioned_users: vec![allowed_user_id.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(), + mentioned_users: vec![], }; let parsed_regex = match mxidwc::parse_pattern("@user.*:example.com") { @@ -106,12 +143,24 @@ async fn process_matrix_messages_in_thread() { 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_basic = super::MatrixMessageProcessingParams::new( + bot_user_id.to_owned(), + Some(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()]); + super::MatrixMessageProcessingParams::new( + bot_user_id.to_owned(), + Some(allowed_users.clone()), + ) + .with_first_message_prefixes_to_strip(vec!["!bai".to_owned()]); + + let message_processing_params_with_bot_user_prefix_stripping = + super::MatrixMessageProcessingParams::new( + bot_user_id.to_owned(), + Some(allowed_users.clone()), + ) + .with_bot_user_prefixes_to_strip(vec!["@baibot: ".to_owned(), "@baibot".to_owned()]); struct TestCase { name: String, @@ -195,10 +244,23 @@ async fn process_matrix_messages_in_thread() { "!bai Hello from the user!".to_owned(), ], }, + TestCase { + name: "Messages that mention the bot user get the bot user prefix stripped" + .to_owned(), + messages: vec![ + allowed_user_message_with_bot_mention.clone(), + allowed_user_message_with_another_user_mention.clone(), + ], + message_processing_params: message_processing_params_with_bot_user_prefix_stripping.clone(), + expected_message_texts: vec![ + "Hello from the user!".to_owned(), + "@baibot: Hello from the user!".to_owned(), + ], + }, ]; for test_case in test_cases { - let processed_messages = super::process_matrix_messages_in_thread( + let processed_messages = super::process_matrix_messages( &test_case.messages, &test_case.message_processing_params, ) @@ -216,3 +278,48 @@ async fn process_matrix_messages_in_thread() { ); } } + +#[test] +fn strip_rich_reply_fallback_text() { + let text = "> <@admin:example.com> What's the difference between Matrix and XMPP?\n\nAnswer me"; + let stripped_text = super::strip_rich_reply_fallback_text(text); + assert_eq!(stripped_text, "Answer me"); +} + +#[test] +fn create_list_of_bot_user_prefixes_to_strip() { + let bot_user_id = + OwnedUserId::try_from("@baibot:example.com").expect("Failed to parse bot user ID"); + + // Test case 1: Bot user with no display name + let bot_display_name = None; + let prefixes = + super::create_list_of_bot_user_prefixes_to_strip(&bot_user_id, &bot_display_name); + + assert_eq!( + prefixes, + vec![ + "@baibot:example.com".to_string(), + "@baibot".to_string(), + "baibot".to_string(), + ":".to_string() + ] + ); + + // Test case 2: Bot user with display name + let bot_display_name = Some("Assistant".to_string()); + let prefixes = + super::create_list_of_bot_user_prefixes_to_strip(&bot_user_id, &bot_display_name); + + assert_eq!( + prefixes, + vec![ + "@baibot:example.com".to_string(), + "@baibot".to_string(), + "baibot".to_string(), + "@Assistant".to_string(), + "Assistant".to_string(), + ":".to_string() + ] + ); +} diff --git a/src/conversation/matrix_llm_bridge.rs b/src/conversation/matrix_llm_bridge.rs index 4528190..522bd2f 100644 --- a/src/conversation/matrix_llm_bridge.rs +++ b/src/conversation/matrix_llm_bridge.rs @@ -1,9 +1,14 @@ +use std::sync::Arc; + use mxlink::matrix_sdk::ruma::OwnedEventId; use mxlink::MatrixLink; +use crate::conversation::matrix::MatrixMessage; + 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, + get_matrix_messages_in_reply_chain, get_matrix_messages_in_thread, process_matrix_messages, + MatrixMessageProcessingParams, RoomEventFetcher, }; pub async fn create_llm_conversation_for_matrix_thread( @@ -14,7 +19,33 @@ pub async fn create_llm_conversation_for_matrix_thread( ) -> Result { 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 llm_messages = filter_messages_and_convert_to_llm_messages(messages, params).await; + + Ok(Conversation { + messages: llm_messages, + }) +} + +pub async fn create_llm_conversation_for_matrix_reply_chain( + event_fetcher: &Arc, + room: &mxlink::matrix_sdk::Room, + event_id: OwnedEventId, + params: &MatrixMessageProcessingParams, +) -> Result { + let messages = get_matrix_messages_in_reply_chain(event_fetcher, room, event_id).await?; + + let llm_messages = filter_messages_and_convert_to_llm_messages(messages, params).await; + + Ok(Conversation { + messages: llm_messages, + }) +} + +async fn filter_messages_and_convert_to_llm_messages( + messages: Vec, + params: &MatrixMessageProcessingParams, +) -> Vec { + let messages_filtered = process_matrix_messages(&messages, params).await; let mut llm_messages: Vec = Vec::new(); @@ -28,7 +59,5 @@ pub async fn create_llm_conversation_for_matrix_thread( llm_messages.push(llm_message); } - Ok(Conversation { - messages: llm_messages, - }) + llm_messages } diff --git a/src/conversation/mod.rs b/src/conversation/mod.rs index 9d02ad3..3ce90da 100644 --- a/src/conversation/mod.rs +++ b/src/conversation/mod.rs @@ -2,4 +2,6 @@ pub(crate) mod llm; pub(crate) mod matrix; mod matrix_llm_bridge; -pub(crate) use matrix_llm_bridge::create_llm_conversation_for_matrix_thread; +pub(crate) use matrix_llm_bridge::{ + create_llm_conversation_for_matrix_reply_chain, create_llm_conversation_for_matrix_thread, +}; diff --git a/src/entity/interaction_context.rs b/src/entity/interaction_context.rs new file mode 100644 index 0000000..3307539 --- /dev/null +++ b/src/entity/interaction_context.rs @@ -0,0 +1,13 @@ +use mxlink::ThreadInfo; + +use super::MessagePayload; + +pub struct InteractionContext { + pub thread_info: ThreadInfo, + pub trigger: InteractionTrigger, +} + +pub struct InteractionTrigger { + pub is_mentioning_bot: bool, + pub payload: MessagePayload, +} diff --git a/src/entity/message_payload.rs b/src/entity/message_payload.rs index 98d09a1..027f69f 100644 --- a/src/entity/message_payload.rs +++ b/src/entity/message_payload.rs @@ -6,10 +6,28 @@ use mxlink::matrix_sdk::ruma::{OwnedEventId, OwnedUserId}; use mxlink::ThreadInfo; /// MessagePayload is like matrix-sdk's MessageType, but represents only message types that the bot deals with and payloads are massaged a bit. +/// +/// This also includes a few synthetic events. #[derive(Debug, Clone)] pub enum MessagePayload { - Text(TextMessageEventContent), + /// A synthetic message payload that indicates that the bot should produce a reply inside a thread. + /// This does not represent an actual message event, it's just a way to trigger a chat completion. + /// + /// When this is invoked, the ThreadInfo contains the full thread details (which represents our context). + /// + /// See: https://github.com/etkecc/baibot/issues/15 + SynthethicChatCompletionTriggerInThread, + /// A synthetic message payload that indicates that the bot should produce a reply to a specific message. + /// This does not represent an actual message event, it's just a way to trigger a chat completion. + /// + /// When this is invoked, the ThreadInfo would refer to the reply-message that triggered us. + /// We can follow the chain upward from it to get the full context. + /// + /// See: https://github.com/etkecc/baibot/issues/15 + SynthethicChatCompletionTriggerForReply, + + Text(TextMessageEventContent), Audio(AudioMessageEventContent), Reaction { diff --git a/src/entity/mod.rs b/src/entity/mod.rs index bb6ed3a..158ef2a 100644 --- a/src/entity/mod.rs +++ b/src/entity/mod.rs @@ -1,15 +1,15 @@ pub mod catch_up_marker; pub mod cfg; pub mod globalconfig; +mod interaction_context; mod message_context; mod message_payload; mod room_config_context; pub mod roomconfig; -mod thread_context; mod trigger_event_info; +pub use interaction_context::{InteractionContext, InteractionTrigger}; pub use message_context::MessageContext; pub use message_payload::MessagePayload; pub use room_config_context::RoomConfigContext; -pub use thread_context::{ThreadContext, ThreadContextFirstMessage}; pub use trigger_event_info::TriggerEventInfo; diff --git a/src/entity/thread_context.rs b/src/entity/thread_context.rs deleted file mode 100644 index 207a18f..0000000 --- a/src/entity/thread_context.rs +++ /dev/null @@ -1,13 +0,0 @@ -use mxlink::ThreadInfo; - -use super::MessagePayload; - -pub struct ThreadContext { - pub info: ThreadInfo, - pub first_message: ThreadContextFirstMessage, -} - -pub struct ThreadContextFirstMessage { - pub is_mentioning_bot: bool, - pub payload: MessagePayload, -}