diff --git a/fluxer_messages/src/shard_impl.rs b/fluxer_messages/src/shard_impl.rs index 7bfe1a52c..78b28ff54 100644 --- a/fluxer_messages/src/shard_impl.rs +++ b/fluxer_messages/src/shard_impl.rs @@ -17,7 +17,7 @@ use crate::types::{ use crate::udt; use chrono::{DateTime, Utc}; use fluxer_svc::shard::ShardService; -use fluxer_svc::transport::NatsTransport; +use fluxer_svc::transport::Transport; use fluxer_svc::{postgres, postgres::BigIntBound, postgres::KeyPart}; use futures::stream::{self, StreamExt}; #[cfg(feature = "scylla")] @@ -81,16 +81,21 @@ const MESSAGE_COLUMNS: &str = "\ has_reaction, version, nsfw_emojis, \ attachments, embeds, sticker_items, message_reference, call, message_snapshots"; -pub struct MessagesShard { +pub struct MessagesShard { storage: MessagesStorage, - transport: NatsTransport, + transport: T, } +#[cfg(test)] +type DeletedMessageKeys = std::sync::Arc>>; + #[derive(Clone)] enum MessagesStorage { Postgres(PostgresMessagesStorage), #[cfg(feature = "scylla")] Scylla(Box), + #[cfg(test)] + Deletions(DeletedMessageKeys), } #[derive(Clone)] @@ -274,8 +279,8 @@ struct MessageMentionContext { embed_users: HashSet, } -impl MessagesShard { - pub fn new_postgres(kv: postgres::KvClient, transport: NatsTransport) -> anyhow::Result { +impl MessagesShard { + pub fn new_postgres(kv: postgres::KvClient, transport: T) -> anyhow::Result { Ok(Self { storage: MessagesStorage::Postgres(PostgresMessagesStorage { kv }), transport, @@ -283,7 +288,7 @@ impl MessagesShard { } #[cfg(feature = "scylla")] - pub async fn new_scylla(db: Arc, transport: NatsTransport) -> anyhow::Result { + pub async fn new_scylla(db: Arc, transport: T) -> anyhow::Result { let stmt_get_by_id = db .prepare(format!( "SELECT {MESSAGE_COLUMNS} FROM messages WHERE channel_id = ? AND bucket = ? AND message_id = ? LIMIT 1" @@ -357,20 +362,8 @@ impl MessagesShard { }) } - fn snowflake_to_bucket(snowflake: i64) -> i32 { - ((snowflake >> 22) / BUCKET_DURATION_MS) as i32 - } - - fn current_bucket() -> i32 { - let now_ms = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as i64; - epoch_millis_to_bucket(now_ms) - } - async fn get_by_id(&self, channel_id: i64, message_id: i64) -> anyhow::Result> { - let bucket = Self::snowflake_to_bucket(message_id); + let bucket = snowflake_to_bucket(message_id); self.storage.get_by_id(channel_id, bucket, message_id).await } @@ -378,7 +371,7 @@ impl MessagesShard { if limit == 0 { return Ok(Vec::new()); } - self.scan_indexed_buckets_desc(channel_id, 0, Self::current_bucket(), limit, None) + self.scan_indexed_buckets_desc(channel_id, 0, current_bucket(), limit, None) .await } @@ -394,7 +387,7 @@ impl MessagesShard { self.scan_indexed_buckets_desc( channel_id, 0, - Self::snowflake_to_bucket(before_id), + snowflake_to_bucket(before_id), limit, Some(before_id), ) @@ -410,8 +403,8 @@ impl MessagesShard { if limit == 0 { return Ok(Vec::new()); } - let min_bucket = Self::snowflake_to_bucket(after_id); - let max_bucket = Self::current_bucket().max(min_bucket); + let min_bucket = snowflake_to_bucket(after_id); + let max_bucket = current_bucket().max(min_bucket); self.scan_indexed_buckets_asc(channel_id, min_bucket, max_bucket, limit, after_id) .await } @@ -672,6 +665,7 @@ impl MessagesShard { return Ok(None); } if message.author_id.is_none() && message.webhook_id.is_none() { + self.cleanup_orphaned_messages(vec![message]).await; return Ok(None); } let context = self @@ -687,11 +681,11 @@ impl MessagesShard { messages: Vec, options: ResponseBuildOptions, ) -> anyhow::Result> { - let messages = messages + let (messages, orphaned_messages): (Vec<_>, Vec<_>) = messages .into_iter() .filter(|message| self.is_message_visible_to_requester(message.message_id, &options)) - .filter(|message| message.author_id.is_some() || message.webhook_id.is_some()) - .collect::>(); + .partition(|message| message.author_id.is_some() || message.webhook_id.is_some()); + self.cleanup_orphaned_messages(orphaned_messages).await; let context = self .build_response_context(&messages, &options, true) .await?; @@ -722,7 +716,7 @@ impl MessagesShard { let count = messages.len(); let failures = stream::iter(messages) .map(|message| async move { - let bucket = Self::snowflake_to_bucket(message.message_id); + let bucket = snowflake_to_bucket(message.message_id); self.storage .delete_message(message.channel_id, bucket, message.message_id) .await @@ -840,7 +834,7 @@ impl MessagesShard { if message.has_reaction == Some(false) { continue; } - let bucket = Self::snowflake_to_bucket(message.message_id); + let bucket = snowflake_to_bucket(message.message_id); groups .entry((message.channel_id, bucket)) .or_default() @@ -1492,6 +1486,8 @@ impl MessagesStorage { MessagesStorage::Scylla(storage) => { storage.get_by_id(channel_id, bucket, message_id).await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(None), } } @@ -1514,6 +1510,8 @@ impl MessagesStorage { .list_buckets_desc(channel_id, min_bucket, max_bucket, limit) .await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(Vec::new()), } } @@ -1536,6 +1534,8 @@ impl MessagesStorage { .list_buckets_asc(channel_id, min_bucket, max_bucket, limit) .await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(Vec::new()), } } @@ -1555,6 +1555,8 @@ impl MessagesStorage { MessagesStorage::Scylla(storage) => { storage.fetch_latest_bucket(channel_id, bucket, limit).await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(Vec::new()), } } @@ -1583,6 +1585,8 @@ impl MessagesStorage { .fetch_before_bucket(channel_id, bucket, before_id, limit) .await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(Vec::new()), } } @@ -1611,6 +1615,8 @@ impl MessagesStorage { .fetch_after_bucket(channel_id, bucket, after_id, limit) .await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(Vec::new()), } } @@ -1628,6 +1634,14 @@ impl MessagesStorage { MessagesStorage::Scylla(storage) => { storage.delete_message(channel_id, bucket, message_id).await } + #[cfg(test)] + MessagesStorage::Deletions(deleted) => { + deleted + .lock() + .expect("deleted message keys mutex poisoned") + .push((channel_id, bucket, message_id)); + Ok(()) + } } } @@ -1660,6 +1674,8 @@ impl MessagesStorage { ) .await } + #[cfg(test)] + MessagesStorage::Deletions(_) => HashMap::new(), } } @@ -1675,6 +1691,8 @@ impl MessagesStorage { MessagesStorage::Scylla(storage) => { storage.fetch_attachment_decay_batch(attachment_ids).await } + #[cfg(test)] + MessagesStorage::Deletions(_) => Ok(HashMap::new()), } } } @@ -2097,7 +2115,7 @@ impl ScyllaMessagesStorage { } } -impl ShardService for MessagesShard { +impl ShardService for MessagesShard { type Request = MessageRequest; type Response = MessageResponse; @@ -2459,6 +2477,18 @@ fn epoch_millis_to_bucket(epoch_millis: i64) -> i32 { ((epoch_millis - FLUXER_EPOCH_MS) / BUCKET_DURATION_MS) as i32 } +fn snowflake_to_bucket(snowflake: i64) -> i32 { + ((snowflake >> 22) / BUCKET_DURATION_MS) as i32 +} + +fn current_bucket() -> i32 { + let now_ms = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as i64; + epoch_millis_to_bucket(now_ms) +} + fn epoch_millis_to_iso(epoch_millis: i64) -> String { DateTime::::from_timestamp_millis(epoch_millis) .unwrap_or(DateTime::::UNIX_EPOCH) @@ -3071,6 +3101,7 @@ impl From for Message { #[cfg(test)] mod tests { use super::*; + use fluxer_svc::transport::InMemoryTransport; use serde_json::json; #[test] @@ -3385,24 +3416,15 @@ mod tests { let message_id = 1_509_197_195_776_110_592; assert_eq!( epoch_millis_to_bucket(snowflake_to_epoch_millis(message_id)), - MessagesShard::snowflake_to_bucket(message_id) + snowflake_to_bucket(message_id) ); } #[test] fn snowflake_bucket_matches_existing_channel_rows() { - assert_eq!( - MessagesShard::snowflake_to_bucket(1_474_193_838_282_432_581), - 406 - ); - assert_eq!( - MessagesShard::snowflake_to_bucket(1_488_946_116_942_273_651), - 410 - ); - assert_eq!( - MessagesShard::snowflake_to_bucket(1_509_256_674_043_502_592), - 416 - ); + assert_eq!(snowflake_to_bucket(1_474_193_838_282_432_581), 406); + assert_eq!(snowflake_to_bucket(1_488_946_116_942_273_651), 410); + assert_eq!(snowflake_to_bucket(1_509_256_674_043_502_592), 416); } #[test] @@ -3515,4 +3537,198 @@ mod tests { HashSet::from([1, 2, 3]) ); } + + fn recording_shard(deleted: &DeletedMessageKeys) -> MessagesShard { + MessagesShard { + storage: MessagesStorage::Deletions(deleted.clone()), + transport: InMemoryTransport::new(), + } + } + + fn build_options() -> ResponseBuildOptions { + ResponseBuildOptions { + viewer_user_id: 1, + source_guild_id: None, + message_history_cutoff_ms: None, + can_read_message_history: true, + media_endpoint: "https://media.example.com".to_owned(), + media_proxy_secret_key: "secret".to_owned(), + include_reactions: false, + nonce: None, + tts: false, + } + } + + fn orphaned_message(message_id: i64) -> Message { + decode_postgres_message(json!({ + "channel_id": {"__fluxer_type": "bigint", "value": "10"}, + "bucket": 416, + "message_id": {"__fluxer_type": "bigint", "value": message_id.to_string()}, + "content": "orphan" + })) + .unwrap() + } + + fn webhook_message(message_id: i64) -> Message { + decode_postgres_message(json!({ + "channel_id": {"__fluxer_type": "bigint", "value": "10"}, + "bucket": 416, + "message_id": {"__fluxer_type": "bigint", "value": message_id.to_string()}, + "webhook_id": {"__fluxer_type": "bigint", "value": "77"}, + "webhook_name": "hook", + "content": "kept" + })) + .unwrap() + } + + fn recorded_deletions(deleted: &DeletedMessageKeys) -> Vec<(i64, i32, i64)> { + deleted.lock().unwrap().clone() + } + + #[tokio::test] + async fn build_path_reaps_orphaned_messages() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let orphan_id = 1_509_197_195_776_110_592; + let kept_id = 1_509_197_195_776_110_593; + + let responses = shard + .build_api_responses_from_messages( + vec![orphaned_message(orphan_id), webhook_message(kept_id)], + build_options(), + ) + .await + .unwrap(); + + assert_eq!(responses.len(), 1); + assert_eq!(responses[0].id, kept_id.to_string()); + assert_eq!( + recorded_deletions(&deleted), + vec![(10, snowflake_to_bucket(orphan_id), orphan_id)] + ); + } + + #[tokio::test] + async fn build_path_reaps_orphans_of_each_batch_separately() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let first_orphan_id = 1_509_197_195_776_110_592; + let second_orphan_id = 1_509_197_195_776_110_594; + + shard + .build_api_responses_from_messages( + vec![ + orphaned_message(first_orphan_id), + webhook_message(1_509_197_195_776_110_593), + ], + build_options(), + ) + .await + .unwrap(); + shard + .build_api_responses_from_messages( + vec![orphaned_message(second_orphan_id)], + build_options(), + ) + .await + .unwrap(); + + assert_eq!( + recorded_deletions(&deleted), + vec![ + (10, snowflake_to_bucket(first_orphan_id), first_orphan_id), + (10, snowflake_to_bucket(second_orphan_id), second_orphan_id), + ] + ); + } + + #[tokio::test] + async fn build_path_reaps_a_repeated_orphan_without_failing() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let orphan_id = 1_509_197_195_776_110_592; + + for _ in 0..2 { + let responses = shard + .build_api_responses_from_messages( + vec![orphaned_message(orphan_id)], + build_options(), + ) + .await + .unwrap(); + assert!(responses.is_empty()); + } + + assert_eq!( + recorded_deletions(&deleted), + vec![ + (10, snowflake_to_bucket(orphan_id), orphan_id), + (10, snowflake_to_bucket(orphan_id), orphan_id), + ] + ); + } + + #[tokio::test] + async fn build_path_leaves_orphans_the_requester_cannot_see() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let orphan_id = 1_509_197_195_776_110_592; + let options = ResponseBuildOptions { + can_read_message_history: false, + message_history_cutoff_ms: Some(snowflake_to_epoch_millis(orphan_id) + 1), + ..build_options() + }; + + let responses = shard + .build_api_responses_from_messages(vec![orphaned_message(orphan_id)], options) + .await + .unwrap(); + + assert!(responses.is_empty()); + assert!(recorded_deletions(&deleted).is_empty()); + } + + #[tokio::test] + async fn single_build_path_reaps_an_orphaned_message() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let orphan_id = 1_509_197_195_776_110_592; + let kept_id = 1_509_197_195_776_110_593; + + let orphan_response = shard + .build_api_response_from_message(orphaned_message(orphan_id), build_options()) + .await + .unwrap(); + let kept_response = shard + .build_api_response_from_message(webhook_message(kept_id), build_options()) + .await + .unwrap(); + + assert!(orphan_response.is_none()); + assert_eq!(kept_response.unwrap().id, kept_id.to_string()); + assert_eq!( + recorded_deletions(&deleted), + vec![(10, snowflake_to_bucket(orphan_id), orphan_id)] + ); + } + + #[tokio::test] + async fn single_build_path_leaves_an_orphan_the_requester_cannot_see() { + let deleted = DeletedMessageKeys::default(); + let shard = recording_shard(&deleted); + let orphan_id = 1_509_197_195_776_110_592; + let options = ResponseBuildOptions { + can_read_message_history: false, + message_history_cutoff_ms: Some(snowflake_to_epoch_millis(orphan_id) + 1), + ..build_options() + }; + + let response = shard + .build_api_response_from_message(orphaned_message(orphan_id), options) + .await + .unwrap(); + + assert!(response.is_none()); + assert!(recorded_deletions(&deleted).is_empty()); + } }