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}