Skip to main content

revolt_database/models/messages/ops/
reference.rs

1use crate::{
2    AppendMessage, FieldsMessage, Message, MessageQuery,
3    PartialMessage, ReferenceDb,
4};
5use futures::future::try_join_all;
6use indexmap::IndexSet;
7use revolt_result::Result;
8use std::collections::HashMap;
9use std::time::SystemTime;
10use ulid::Ulid;
11
12use super::AbstractMessages;
13
14#[async_trait]
15impl AbstractMessages for ReferenceDb {
16    /// Insert a new message into the database
17    async fn insert_message(&self, message: &Message) -> Result<()> {
18        let mut messages = self.messages.lock().await;
19        if messages.contains_key(&message.id) {
20            Err(create_database_error!("insert", "message"))
21        } else {
22            messages.insert(message.id.to_string(), message.clone());
23            Ok(())
24        }
25    }
26
27    /// Fetch a message by its id
28    async fn fetch_message(&self, id: &str) -> Result<Message> {
29        let messages = self.messages.lock().await;
30        messages
31            .get(id)
32            .cloned()
33            .ok_or_else(|| create_error!(NotFound))
34    }
35
36    /// Fetch multiple messages by given query
37    async fn fetch_messages(&self, query: MessageQuery) -> Result<Vec<Message>> {
38        let messages = self.messages.lock().await;
39        let matched_messages = messages
40            .values()
41            .filter(|message| {
42                if let Some(channel) = &query.filter.channel {
43                    if &message.channel != channel {
44                        return false;
45                    }
46                }
47
48                if let Some(author) = &query.filter.author {
49                    if &message.author != author {
50                        return false;
51                    }
52                }
53
54                if let Some(query) = &query.filter.query {
55                    if let Some(content) = &message.content {
56                        if !content.to_lowercase().contains(query) {
57                            return false;
58                        }
59                    } else {
60                        return false;
61                    }
62                }
63
64                if let Some(pinned) = query.filter.pinned {
65                    if message.pinned.unwrap_or_default() == pinned {
66                        return false;
67                    }
68                }
69
70                true
71            })
72            .cloned()
73            .collect();
74
75        // FIXME: sorting, etc (will be required for tests)
76
77        Ok(matched_messages)
78
79        /*
80        // 2. Find query limit
81        let limit = query.limit.unwrap_or(50);
82
83        // 3. Apply message time period
84        match query.time_period {
85            MessageTimePeriod::Relative { nearby } => {
86                // 3.1. Prepare filters
87                let mut older_message_filter = filter.clone();
88                let mut newer_message_filter = filter;
89
90                older_message_filter.insert(
91                    "_id",
92                    doc! {
93                        "$lt": &nearby
94                    },
95                );
96
97                newer_message_filter.insert(
98                    "_id",
99                    doc! {
100                        "$gte": &nearby
101                    },
102                );
103
104                // 3.2. Execute in both directions
105                let (a, b) = try_join!(
106                    self.find_with_options::<_, Message>(
107                        COL,
108                        newer_message_filter,
109                        FindOptions::builder()
110                            .limit(limit / 2 + 1)
111                            .sort(doc! {
112                                "_id": 1_i32
113                            })
114                            .build(),
115                    ),
116                    self.find_with_options::<_, Message>(
117                        COL,
118                        older_message_filter,
119                        FindOptions::builder()
120                            .limit(limit / 2)
121                            .sort(doc! {
122                                "_id": -1_i32
123                            })
124                            .build(),
125                    )
126                )
127                .map_err(|_| create_database_error!("find", COL))?;
128
129                Ok([a, b].concat())
130            }
131            MessageTimePeriod::Absolute {
132                before,
133                after,
134                sort,
135            } => {
136                // 3.1. Apply message ID filter
137                if let Some(doc) = match (before, after) {
138                    (Some(before), Some(after)) => Some(doc! {
139                        "$lt": before,
140                        "$gt": after
141                    }),
142                    (Some(before), _) => Some(doc! {
143                        "$lt": before
144                    }),
145                    (_, Some(after)) => Some(doc! {
146                        "$gt": after
147                    }),
148                    _ => None,
149                } {
150                    filter.insert("_id", doc);
151                }
152
153                // 3.2. Execute with given message sort
154                self.find_with_options(
155                    COL,
156                    filter,
157                    FindOptions::builder()
158                        .limit(limit)
159                        .sort(match sort.unwrap_or(MessageSort::Latest) {
160                            // Sort by relevance, fallback to latest
161                            MessageSort::Relevance => {
162                                if is_search_query {
163                                    doc! {
164                                        "score": {
165                                            "$meta": "textScore"
166                                        }
167                                    }
168                                } else {
169                                    doc! {
170                                        "_id": -1_i32
171                                    }
172                                }
173                            }
174                            // Sort by latest first
175                            MessageSort::Latest => doc! {
176                                "_id": -1_i32
177                            },
178                            // Sort by oldest first
179                            MessageSort::Oldest => doc! {
180                                "_id": 1_i32
181                            },
182                        })
183                        .build(),
184                )
185                .await
186                .map_err(|_| create_database_error!("find", COL))
187            }
188        }*/
189    }
190
191    /// Fetch multiple messages by given IDs
192    async fn fetch_messages_by_id(&self, ids: &[String]) -> Result<Vec<Message>> {
193        try_join_all(ids.iter().map(|id| self.fetch_message(id))).await
194    }
195
196    /// Update a given message with new information
197    async fn update_message(
198        &self,
199        id: &str,
200        message: &PartialMessage,
201        remove: Vec<FieldsMessage>,
202    ) -> Result<()> {
203        let mut messages = self.messages.lock().await;
204        if let Some(message_data) = messages.get_mut(id) {
205            message_data.apply_options(message.to_owned());
206
207            for field in remove {
208                #[allow(clippy::disallowed_methods)]
209                message_data.remove_field(&field);
210            }
211            Ok(())
212        } else {
213            Err(create_error!(NotFound))
214        }
215    }
216
217    /// Append information to a given message
218    async fn append_message(&self, id: &str, append: &AppendMessage) -> Result<()> {
219        let mut messages = self.messages.lock().await;
220        if let Some(message_data) = messages.get_mut(id) {
221            if let Some(embeds) = &append.embeds {
222                if !embeds.is_empty() {
223                    if let Some(embeds_data) = &mut message_data.embeds {
224                        embeds_data.extend(embeds.clone());
225                    } else {
226                        message_data.embeds = Some(embeds.clone());
227                    }
228                }
229            }
230
231            Ok(())
232        } else {
233            Err(create_error!(NotFound))
234        }
235    }
236
237    /// Add a new reaction to a message
238    async fn add_reaction(&self, id: &str, emoji: &str, user: &str) -> Result<()> {
239        let mut messages = self.messages.lock().await;
240        if let Some(message) = messages.get_mut(id) {
241            if let Some(users) = message.reactions.get_mut(emoji) {
242                users.insert(user.to_string());
243            } else {
244                message
245                    .reactions
246                    .insert(emoji.to_string(), IndexSet::from([user.to_string()]));
247            }
248
249            Ok(())
250        } else {
251            Err(create_error!(NotFound))
252        }
253    }
254
255    /// Remove a reaction from a message
256    async fn remove_reaction(&self, id: &str, emoji: &str, user: &str) -> Result<()> {
257        let mut messages = self.messages.lock().await;
258        if let Some(message) = messages.get_mut(id) {
259            if let Some(users) = message.reactions.get_mut(emoji) {
260                users.swap_remove(&user.to_string());
261            }
262
263            Ok(())
264        } else {
265            Err(create_error!(NotFound))
266        }
267    }
268
269    /// Remove reaction from a message
270    async fn clear_reaction(&self, id: &str, emoji: &str) -> Result<()> {
271        let mut messages = self.messages.lock().await;
272        if let Some(message) = messages.get_mut(id) {
273            message.reactions.swap_remove(emoji);
274            Ok(())
275        } else {
276            Err(create_error!(NotFound))
277        }
278    }
279
280    /// Delete a message from the database by its id
281    async fn delete_message(&self, id: &str) -> Result<()> {
282        let mut messages = self.messages.lock().await;
283        if messages.remove(id).is_some() {
284            Ok(())
285        } else {
286            Err(create_error!(NotFound))
287        }
288    }
289
290    /// Delete messages from a channel by their ids and corresponding channel id
291    async fn delete_messages(&self, channel: &str, ids: &[String]) -> Result<()> {
292        self.messages
293            .lock()
294            .await
295            .retain(|id, message| message.channel != channel && !ids.contains(id));
296
297        Ok(())
298    }
299
300    /// Delete all messages from a specific author in a list of channels from a certain ULID onwards
301    async fn delete_messages_by_author_since(
302        &self,
303        channels: &[String],
304        author: &str,
305        since: SystemTime,
306    ) -> Result<HashMap<String, Vec<String>>> {
307        let threshold_ulid = Ulid::from_datetime(since).to_string();
308        let mut deleted_messages: HashMap<String, Vec<String>> = HashMap::new();
309        let mut attachment_ids: Vec<String> = Vec::new();
310
311        let messages = self.messages.lock().await;
312
313        // First pass: collect attachment IDs and message IDs to delete
314        for (id, message) in messages.iter() {
315            let should_delete = message.author == author
316                && channels.contains(&message.channel)
317                && id.as_str() >= threshold_ulid.as_str();
318
319            if should_delete {
320                // Collect attachment IDs
321                if let Some(attachments) = &message.attachments {
322                    for attachment in attachments {
323                        attachment_ids.push(attachment.id.clone());
324                    }
325                }
326
327                deleted_messages
328                    .entry(message.channel.clone())
329                    .or_default()
330                    .push(id.clone());
331            }
332        }
333        drop(messages);
334
335        // Mark attachments as deleted
336        if !attachment_ids.is_empty() {
337            let mut files = self.files.lock().await;
338            for attachment_id in attachment_ids {
339                if let Some(file) = files.get_mut(&attachment_id) {
340                    file.deleted = Some(true);
341                }
342            }
343        }
344
345        // Delete the messages
346        self.messages.lock().await.retain(|id, message| {
347            let should_keep = !(message.author == author
348                && channels.contains(&message.channel)
349                && id.as_str() >= threshold_ulid.as_str());
350            should_keep
351        });
352
353        Ok(deleted_messages)
354    }
355
356    async fn delete_messages_by_user(&self, user_id: &str) -> Result<()> {
357        let mut messages = self.messages.lock().await;
358
359        messages.retain(|_, message| message.author != user_id);
360
361        // TODO: remove attachments as well
362
363        Ok(())
364    }
365}