Skip to main content

revolt_database/models/channels/ops/
reference.rs

1use std::collections::hash_map::Entry;
2
3use super::AbstractChannels;
4use crate::ReferenceDb;
5use crate::util::ChunkedDatabaseGenerator;
6use crate::{Channel, FieldsChannel, PartialChannel};
7use revolt_permissions::OverrideField;
8use revolt_result::Result;
9
10#[async_trait]
11impl AbstractChannels for ReferenceDb {
12    /// Insert a new channel in the database
13    async fn insert_channel(&self, channel: &Channel) -> Result<()> {
14        let mut channels = self.channels.lock().await;
15        if let Entry::Vacant(entry) = channels.entry(channel.id().to_string()) {
16            entry.insert(channel.clone());
17            Ok(())
18        } else {
19            Err(create_database_error!("insert", "channel"))
20        }
21    }
22
23    /// Fetch a channel from the database
24    async fn fetch_channel(&self, channel_id: &str) -> Result<Channel> {
25        let channels = self.channels.lock().await;
26        channels
27            .get(channel_id)
28            .cloned()
29            .ok_or_else(|| create_error!(NotFound))
30    }
31
32    /// Fetch all channels from the database
33    async fn fetch_channels<'a>(&self, ids: &'a [String]) -> Result<Vec<Channel>> {
34        let channels = self.channels.lock().await;
35        ids.iter()
36            .map(|id| {
37                channels
38                    .get(id)
39                    .cloned()
40                    .ok_or_else(|| create_error!(NotFound))
41            })
42            .collect()
43    }
44
45    /// Fetch all direct messages for a user
46    async fn find_direct_messages(&self, user_id: &str) -> Result<Vec<Channel>> {
47        let channels = self.channels.lock().await;
48        Ok(channels
49            .values()
50            .filter(|channel| channel.contains_user(user_id))
51            .cloned()
52            .collect())
53    }
54
55    // Fetch all group dms for a user
56    async fn find_group_message_channels(&self, user_id: &str) -> Result<ChunkedDatabaseGenerator<Channel>> {
57        let channels = self.channels.lock().await;
58        let groups = channels
59            .values()
60            .filter(|channel| match channel {
61                Channel::Group { recipients, .. } => {
62                    recipients.iter().any(|recipient| recipient == user_id)
63                }
64                _ => false,
65            })
66            .cloned()
67            .collect();
68
69        Ok(ChunkedDatabaseGenerator::new_reference(groups))
70    }
71
72    // Fetch saved messages channel
73    async fn find_saved_messages_channel(&self, user_id: &str) -> Result<Channel> {
74        let channels = self.channels.lock().await;
75        channels
76            .get(user_id)
77            .cloned()
78            .ok_or_else(|| create_database_error!("fetch", "channel"))
79    }
80
81    // Fetch direct message channel (DM or Saved Messages)
82    async fn find_direct_message_channel(&self, user_a: &str, user_b: &str) -> Result<Channel> {
83        let channels = self.channels.lock().await;
84        for (_, data) in channels.iter() {
85            if data.contains_user(user_a) && data.contains_user(user_b) {
86                return Ok(data.to_owned());
87            }
88        }
89        Err(create_error!(NotFound))
90    }
91    /// Insert a user to a group
92    async fn add_user_to_group(&self, channel_id: &str, user_id: &str) -> Result<()> {
93        let mut channels = self.channels.lock().await;
94
95        if let Some(Channel::Group { recipients, .. }) = channels.get_mut(channel_id) {
96            recipients.push(String::from(user_id));
97            Ok(())
98        } else {
99            Err(create_error!(InvalidOperation))
100        }
101    }
102    /// Insert channel role permissions
103    async fn set_channel_role_permission(
104        &self,
105        channel_id: &str,
106        role_id: &str,
107        permissions: OverrideField,
108    ) -> Result<()> {
109        let mut channels = self.channels.lock().await;
110
111        if let Some(mut channel) = channels.get_mut(channel_id) {
112            match &mut channel {
113                Channel::TextChannel {
114                    role_permissions, ..
115                } => {
116                    if role_permissions.get(role_id).is_some() {
117                        role_permissions.remove(role_id);
118                        role_permissions.insert(String::from(role_id), permissions);
119
120                        Ok(())
121                    } else {
122                        Err(create_error!(NotFound))
123                    }
124                }
125                _ => Err(create_error!(NotFound)),
126            }
127        } else {
128            Err(create_error!(NotFound))
129        }
130    }
131
132    // Update channel
133    async fn update_channel(
134        &self,
135        id: &str,
136        channel: &PartialChannel,
137        remove: Vec<FieldsChannel>,
138    ) -> Result<()> {
139        let mut channels = self.channels.lock().await;
140        if let Some(channel_data) = channels.get_mut(id) {
141            channel_data.apply_options(channel.to_owned());
142            channel_data.remove_fields(remove);
143            Ok(())
144        } else {
145            Err(create_error!(NotFound))
146        }
147    }
148
149    // Remove a user from a group
150    async fn remove_user_from_group(&self, channel: &str, user: &str) -> Result<()> {
151        let mut channels = self.channels.lock().await;
152        if let Some(Channel::Group { recipients, .. }) = channels.get_mut(channel) {
153            if let Some(index) = recipients.iter().position(|recipient| recipient == user) {
154                recipients.remove(index);
155                return Ok(());
156            } else {
157                return Err(create_error!(NotFound));
158            }
159        }
160        Err(create_error!(NotFound))
161    }
162
163    // Remove a user from all specified groups
164    async fn remove_user_from_groups(&self, channel_ids: Vec<String>, user_id: &str) -> Result<()> {
165        let mut channels = self.channels.lock().await;
166
167        for channel_id in channel_ids {
168            if let Some(Channel::Group { recipients, .. }) = channels.get_mut(&channel_id) {
169                recipients.retain(|recipient| recipient != user_id);
170            }
171        };
172
173        Ok(())
174    }
175
176    // Delete a channel
177    async fn delete_channel(&self, channel: &Channel) -> Result<()> {
178        let mut channels = self.channels.lock().await;
179        if channels.remove(channel.id()).is_some() {
180            Ok(())
181        } else {
182            Err(create_error!(NotFound))
183        }
184    }
185}