diff --git a/fluxer_messages/src/mention_extractor.rs b/fluxer_messages/src/mention_extractor.rs index ad6d316fe..bcccf29c8 100644 --- a/fluxer_messages/src/mention_extractor.rs +++ b/fluxer_messages/src/mention_extractor.rs @@ -32,9 +32,25 @@ pub fn extract_mentions_from_markdown(input: Option<&str>) -> MessageMentions { let Some(input) = input else { return MessageMentions::default(); }; - if input.is_empty() { + if input.is_empty() || !may_contain_mention(input) { return MessageMentions::default(); } + parse_mentions(input) +} + +fn may_contain_mention(input: &str) -> bool { + let bytes = input.as_bytes(); + bytes.iter().enumerate().any(|(index, byte)| match byte { + b'<' => matches!(bytes.get(index + 1), Some(b'@' | b'#')), + b'@' => { + let rest = &bytes[index..]; + rest.starts_with(b"@everyone") || rest.starts_with(b"@here") + } + _ => false, + }) +} + +fn parse_mentions(input: &str) -> MessageMentions { let cleaned = blank_raw_urls(input); let mut parser = MarkdownParser::new(ParserFlags::ALL, EmojiContext::default()); let Ok(nodes) = parser.parse(cleaned.as_ref()) else { @@ -149,7 +165,94 @@ fn url_finder() -> &'static LinkFinder { #[cfg(test)] mod tests { - use super::extract_mentions_from_markdown; + use super::{extract_mentions_from_markdown, may_contain_mention, parse_mentions}; + + const MENTION_CORPUS: &[&str] = &[ + "hey, are we still on for tonight?", + "lol", + "that build is green now, shipping it", + "no idea, ask in the other channel", + "brb", + "https://example.com/some/long/path?query=1&other=2", + "check out https://github.com/fluxerapp/fluxer/pull/1234 when you get a sec", + "mail me at someone@example.com", + "the price is 12 <> 15 depending on the region", + "a < b && c > d", + "", + "", + "<+15550001111>", + "", + "", + "<:custom_emoji:1234567890>", + "", + "```rust\nfn main() { println!(\"hi\"); }\n```", + "`inline code with a # and an < in it`", + "**bold** *italic* __underline__ ~~strike~~ ||spoiler||", + "> quoted line\n> another quoted line", + "# heading\n## smaller heading\n-# subtext", + "- item one\n- item two\n1. numbered", + "|a|b|\n|-|-|\n|1|2|", + "[masked link](https://example.com)", + "emoji party πŸŽ‰πŸŽ‰πŸŽ‰ and a flag πŸ‡ΈπŸ‡ͺ", + "escaped \\<@123> should stay text", + "escaped \\@everyone should stay text", + "channel #general is over there", + "email me @ work tomorrow", + "@ everyone with a space", + "@ here with a space", + "@everyones and @heresy are longer words", + "hi <@123> <@!456> <@&789> <#321>", + "<@0> <@&0> <#0>", + "@everyone hello @here", + "`@everyone` @here\n```\n@everyone\n```", + "`<@111>` <@222>\n```txt\n<@333> <#444>\n```\n<#555>", + "https://example.com/<@123> <@456>", + "[click here](https://example.com) and <#888> <@999>", + "**bold <@100>** *italic <@&200>* ~~strike <#300>~~ __underline <@400>__", + "> quoted <@111>\n<@222>", + "||spoiler <@333>||", + "[<@101>](https://example.com/<@202>) <#303>", + "<@123456789012345678> <@&999> <#888>", + "<#not_an_id> <@not_an_id> <@¬_an_id>", + ]; + + #[test] + fn prefilter_never_changes_extraction_over_the_corpus() { + for input in MENTION_CORPUS { + assert_eq!( + extract_mentions_from_markdown(Some(input)), + parse_mentions(input), + "prefilter changed the result for {input:?}" + ); + } + } + + #[test] + fn corpus_exercises_both_prefilter_outcomes() { + assert!( + MENTION_CORPUS + .iter() + .any(|input| may_contain_mention(input)) + ); + assert!( + MENTION_CORPUS + .iter() + .any(|input| !may_contain_mention(input)) + ); + } + + #[test] + fn prefilter_accepts_every_mention_marker() { + assert!(may_contain_mention("<@1>")); + assert!(may_contain_mention("<@!1>")); + assert!(may_contain_mention("<@&1>")); + assert!(may_contain_mention("<#1>")); + assert!(may_contain_mention("@everyone")); + assert!(may_contain_mention("@here")); + assert!(!may_contain_mention("< @1>")); + assert!(!may_contain_mention("@ everyone")); + assert!(!may_contain_mention("plain text")); + } #[test] fn extracts_real_user_role_and_channel_mentions() { diff --git a/fluxer_messages/src/shard_impl.rs b/fluxer_messages/src/shard_impl.rs index 09d10791c..d2dd69ca4 100644 --- a/fluxer_messages/src/shard_impl.rs +++ b/fluxer_messages/src/shard_impl.rs @@ -271,6 +271,7 @@ struct ResponseContext { struct MessageMentionContext { content: MessageMentions, snapshots: Vec, + embed_users: HashSet, } impl MessagesShard { @@ -1051,11 +1052,15 @@ impl MessagesShard { .iter() .filter_map(map_sticker) .collect(); - let content_mentions = context - .mention_context - .get(&message.message_id) - .map(|mentions| mentions.content.clone()) - .unwrap_or_else(|| extract_mentions_from_markdown(message.content.as_deref())); + let fallback_mentions; + let message_mentions = match context.mention_context.get(&message.message_id) { + Some(mentions) => mentions, + None => { + fallback_mentions = build_mention_context_entry(message); + &fallback_mentions + } + }; + let content_mentions = &message_mentions.content; let mention_roles = ids_present_in_set(&message.mention_roles, &content_mentions.roles); let mention_channels = ids_present_in_set(&message.mention_channels, &content_mentions.channels) @@ -1063,24 +1068,9 @@ impl MessagesShard { .filter_map(|id| context.channel_mentions.get(&id).cloned()) .collect::>(); let mut referenced_user_ids = content_mentions.users.clone(); - for embed in message.embeds.as_deref().unwrap_or_default() { - collect_user_ids_from_embed(embed, &mut referenced_user_ids); - } - if let Some(snapshots) = &message.message_snapshots { - for (index, snapshot) in snapshots.iter().enumerate() { - let snapshot_mentions = context - .mention_context - .get(&message.message_id) - .and_then(|mentions| mentions.snapshots.get(index)) - .cloned() - .unwrap_or_else(|| extract_mentions_from_markdown(snapshot.content.as_deref())); - referenced_user_ids.extend(snapshot_mentions.users); - if let Some(embeds) = &snapshot.embeds { - for embed in embeds { - collect_user_ids_from_embed(embed, &mut referenced_user_ids); - } - } - } + referenced_user_ids.extend(message_mentions.embed_users.iter().copied()); + for snapshot_mentions in &message_mentions.snapshots { + referenced_user_ids.extend(snapshot_mentions.users.iter().copied()); } let mentioned_user_ids = message .mention_users @@ -2486,25 +2476,32 @@ fn map_embed_field_response(field: MessageEmbedField) -> ApiEmbedFieldResponse { fn build_message_mention_context(messages: &[&Message]) -> HashMap { messages .iter() - .map(|message| { - let snapshots = message - .message_snapshots - .as_deref() - .unwrap_or_default() - .iter() - .map(|snapshot| extract_mentions_from_markdown(snapshot.content.as_deref())) - .collect(); - ( - message.message_id, - MessageMentionContext { - content: extract_mentions_from_markdown(message.content.as_deref()), - snapshots, - }, - ) - }) + .map(|message| (message.message_id, build_mention_context_entry(message))) .collect() } +fn build_mention_context_entry(message: &Message) -> MessageMentionContext { + let message_snapshots = message.message_snapshots.as_deref().unwrap_or_default(); + let snapshots = message_snapshots + .iter() + .map(|snapshot| extract_mentions_from_markdown(snapshot.content.as_deref())) + .collect(); + let mut embed_users = HashSet::new(); + for embed in message.embeds.as_deref().unwrap_or_default() { + collect_user_ids_from_embed(embed, &mut embed_users); + } + for snapshot in message_snapshots { + for embed in snapshot.embeds.as_deref().unwrap_or_default() { + collect_user_ids_from_embed(embed, &mut embed_users); + } + } + MessageMentionContext { + content: extract_mentions_from_markdown(message.content.as_deref()), + snapshots, + embed_users, + } +} + fn ids_present_in_set(ids: &[i64], present: &HashSet) -> Vec { ids.iter() .filter(|id| present.contains(id)) @@ -2581,11 +2578,7 @@ fn collect_user_ids( } if let Some(mentions) = mention_context.get(&message.message_id) { ids.extend(mentions.content.users.iter().copied()); - } - if let Some(embeds) = &message.embeds { - for embed in embeds { - collect_user_ids_from_embed(embed, &mut ids); - } + ids.extend(mentions.embed_users.iter().copied()); } if let Some(snapshots) = &message.message_snapshots { for (index, snapshot) in snapshots.iter().enumerate() { @@ -2598,11 +2591,6 @@ fn collect_user_ids( { ids.extend(mentions.users.iter().copied()); } - if let Some(embeds) = &snapshot.embeds { - for embed in embeds { - collect_user_ids_from_embed(embed, &mut ids); - } - } } } } @@ -3224,6 +3212,43 @@ mod tests { ); } + #[test] + fn mention_context_carries_embed_user_ids_for_message_and_snapshots() { + let message: Message = serde_json::from_value(json!({ + "message_id": "10", + "channel_id": "20", + "bucket": 1, + "author_id": "30", + "type": 0, + "version": 0, + "content": "hello <@40>", + "mention_users": ["40"], + "embeds": [{ + "title": "title <@50>", + "description": "description <@60>", + "footer": {"text": "footer <@70>"}, + "fields": [{"name": "field <@80>", "value": "value <@90>"}] + }], + "message_snapshots": [{ + "content": "snapshot <@100>", + "embeds": [{"description": "snapshot embed <@110>"}] + }] + })) + .unwrap(); + let messages = std::slice::from_ref(&message); + + let mention_context = build_message_mention_context(messages); + let entry = mention_context.get(&10).unwrap(); + + assert_eq!(entry.content.users, HashSet::from([40])); + assert_eq!(entry.embed_users, HashSet::from([50, 60, 70, 80, 90, 110])); + assert_eq!(entry.snapshots[0].users, HashSet::from([100])); + assert_eq!( + collect_user_ids(messages, &mention_context), + HashSet::from([30, 40, 50, 60, 70, 80, 90, 100, 110]) + ); + } + #[test] fn postgres_reaction_decoder_maps_created_at() { let (message_id, reaction) = decode_postgres_reaction(json!({