Skip to main content

revolt_database/models/servers/
model.rs

1use std::collections::{HashMap, HashSet};
2
3use revolt_models::v0::{self, DataCreateServerChannel};
4use revolt_permissions::{OverrideField, DEFAULT_PERMISSION_SERVER};
5use revolt_result::Result;
6use ulid::Ulid;
7
8use crate::{events::client::EventV1, Channel, Database, File, User};
9
10auto_derived_partial!(
11    /// Server
12    pub struct Server {
13        /// Unique Id
14        #[serde(rename = "_id")]
15        pub id: String,
16        /// User id of the owner
17        pub owner: String,
18
19        /// Name of the server
20        pub name: String,
21        /// Description for the server
22        #[serde(skip_serializing_if = "Option::is_none")]
23        pub description: Option<String>,
24
25        /// Channels within this server
26        // TODO: investigate if this is redundant and can be removed
27        pub channels: Vec<String>,
28        /// Categories for this server
29        #[serde(skip_serializing_if = "Option::is_none")]
30        pub categories: Option<Vec<Category>>,
31        /// Configuration for sending system event messages
32        #[serde(skip_serializing_if = "Option::is_none")]
33        pub system_messages: Option<SystemMessageChannels>,
34
35        /// Roles for this server
36        #[serde(
37            default = "HashMap::<String, Role>::new",
38            skip_serializing_if = "HashMap::<String, Role>::is_empty"
39        )]
40        pub roles: HashMap<String, Role>,
41        /// Default set of server and channel permissions
42        pub default_permissions: i64,
43
44        /// Icon attachment
45        #[serde(skip_serializing_if = "Option::is_none")]
46        pub icon: Option<File>,
47        /// Banner attachment
48        #[serde(skip_serializing_if = "Option::is_none")]
49        pub banner: Option<File>,
50
51        /// Bitfield of server flags
52        #[serde(skip_serializing_if = "Option::is_none")]
53        pub flags: Option<i32>,
54
55        /// Whether this server is flagged as not safe for work
56        #[serde(skip_serializing_if = "crate::if_false", default)]
57        pub nsfw: bool,
58        /// Whether to enable analytics
59        #[serde(skip_serializing_if = "crate::if_false", default)]
60        pub analytics: bool,
61        /// Whether this server should be publicly discoverable
62        #[serde(skip_serializing_if = "crate::if_false", default)]
63        pub discoverable: bool,
64    },
65    "PartialServer"
66);
67
68auto_derived_partial!(
69    /// Role
70    pub struct Role {
71        /// Unique Id
72        #[serde(rename = "_id")]
73        pub id: String,
74        /// Role name
75        pub name: String,
76        /// Permissions available to this role
77        pub permissions: OverrideField,
78        /// Colour used for this role
79        ///
80        /// This can be any valid CSS colour
81        #[serde(skip_serializing_if = "Option::is_none")]
82        pub colour: Option<String>,
83        /// Whether this role should be shown separately on the member sidebar
84        #[serde(skip_serializing_if = "crate::if_false", default)]
85        pub hoist: bool,
86        /// Ranking of this role
87        #[serde(default)]
88        pub rank: i64,
89        /// Custom icon attachment
90        #[serde(skip_serializing_if = "Option::is_none")]
91        pub icon: Option<File>,
92    },
93    "PartialRole"
94);
95
96auto_derived!(
97    /// Channel category
98    pub struct Category {
99        /// Unique ID for this category
100        pub id: String,
101        /// Title for this category
102        pub title: String,
103        /// Channels in this category
104        pub channels: Vec<String>,
105    }
106
107    /// System message channel assignments
108    pub struct SystemMessageChannels {
109        /// ID of channel to send user join messages in
110        #[serde(skip_serializing_if = "Option::is_none")]
111        pub user_joined: Option<String>,
112        /// ID of channel to send user left messages in
113        #[serde(skip_serializing_if = "Option::is_none")]
114        pub user_left: Option<String>,
115        /// ID of channel to send user kicked messages in
116        #[serde(skip_serializing_if = "Option::is_none")]
117        pub user_kicked: Option<String>,
118        /// ID of channel to send user banned messages in
119        #[serde(skip_serializing_if = "Option::is_none")]
120        pub user_banned: Option<String>,
121    }
122
123    /// Optional fields on server object
124    pub enum FieldsServer {
125        Description,
126        Categories,
127        SystemMessages,
128        Icon,
129        Banner,
130    }
131
132    /// Optional fields on server object
133    pub enum FieldsRole {
134        Colour,
135        Icon,
136    }
137);
138
139#[allow(clippy::disallowed_methods)]
140impl Server {
141    /// Create a server
142    pub async fn create(
143        db: &Database,
144        data: v0::DataCreateServer,
145        owner: &User,
146        create_default_channels: bool,
147    ) -> Result<(Server, Vec<Channel>)> {
148        let mut server = Server {
149            id: ulid::Ulid::new().to_string(),
150            owner: owner.id.to_string(),
151            name: data.name,
152            description: data.description,
153            channels: vec![],
154            nsfw: data.nsfw.unwrap_or(false),
155            default_permissions: *DEFAULT_PERMISSION_SERVER as i64,
156
157            analytics: false,
158            banner: None,
159            categories: None,
160            discoverable: false,
161            flags: None,
162            icon: None,
163            roles: HashMap::new(),
164            system_messages: None,
165        };
166
167        let channels: Vec<Channel> = if create_default_channels {
168            vec![
169                Channel::create_server_channel(
170                    db,
171                    &mut server,
172                    DataCreateServerChannel {
173                        channel_type: v0::LegacyServerChannelType::Text,
174                        name: "General".to_string(),
175                        ..Default::default()
176                    },
177                    false,
178                )
179                .await?,
180            ]
181        } else {
182            vec![]
183        };
184
185        server.channels = channels.iter().map(|c| c.id().to_string()).collect();
186        db.insert_server(&server).await?;
187        Ok((server, channels))
188    }
189
190    /// Update server data
191    pub async fn update(
192        &mut self,
193        db: &Database,
194        partial: PartialServer,
195        remove: Vec<FieldsServer>,
196    ) -> Result<()> {
197        for field in &remove {
198            self.remove_field(field);
199        }
200
201        self.apply_options(partial.clone());
202
203        db.update_server(&self.id, &partial, remove.clone()).await?;
204
205        EventV1::ServerUpdate {
206            id: self.id.clone(),
207            data: partial.into(),
208            clear: remove.into_iter().map(|v| v.into()).collect(),
209        }
210        .p(self.id.clone())
211        .await;
212
213        Ok(())
214    }
215
216    /// Delete a server
217    pub async fn delete(self, db: &Database) -> Result<()> {
218        EventV1::ServerDelete {
219            id: self.id.clone(),
220        }
221        .p(self.id.clone())
222        .await;
223
224        db.delete_server(&self.id).await
225    }
226
227    /// Remove a field from Server
228    pub fn remove_field(&mut self, field: &FieldsServer) {
229        match field {
230            FieldsServer::Description => self.description = None,
231            FieldsServer::Categories => self.categories = None,
232            FieldsServer::SystemMessages => self.system_messages = None,
233            FieldsServer::Icon => self.icon = None,
234            FieldsServer::Banner => self.banner = None,
235        }
236    }
237
238    /// Generates a PartialServer containing the data which has changed in an update
239    pub fn generate_diff(&self, partial: &PartialServer, remove: &[FieldsServer]) -> PartialServer {
240        let mut before = PartialServer::default();
241
242        generate_diff!(
243            self, before, partial, remove,
244            (
245                owner,
246                name,
247                (FieldsServer::Description) description,
248                (FieldsServer::Categories) categories,
249                (FieldsServer::SystemMessages) system_messages,
250                roles,
251                default_permissions,
252                (FieldsServer::Icon) icon,
253                (FieldsServer::Banner) banner,
254                nsfw,
255                analytics,
256                discoverable,
257            )
258        );
259
260        before
261    }
262
263    /// Ordered roles list
264    pub fn ordered_roles(&self) -> Vec<(String, Role)> {
265        let mut ordered_roles = self.roles.clone().into_iter().collect::<Vec<_>>();
266        ordered_roles.sort_by(|(_, role_a), (_, role_b)| role_a.rank.cmp(&role_b.rank));
267        ordered_roles
268    }
269
270    /// Set role permission on a server
271    pub async fn set_role_permission(
272        &mut self,
273        db: &Database,
274        role_id: &str,
275        permissions: OverrideField,
276    ) -> Result<()> {
277        if let Some(role) = self.roles.get_mut(role_id) {
278            role.update(
279                db,
280                &self.id,
281                PartialRole {
282                    permissions: Some(permissions),
283                    ..Default::default()
284                },
285                vec![],
286            )
287            .await?;
288
289            Ok(())
290        } else {
291            Err(create_error!(NotFound))
292        }
293    }
294
295    /// Reorders the server's roles rankings
296    pub async fn set_role_ordering(&mut self, db: &Database, new_order: Vec<String>) -> Result<()> {
297        // New order must always contain every role
298        debug_assert_eq!(self.roles.len(), new_order.len());
299
300        // Set the role's ranks to the positions in the vec
301        for (rank, id) in new_order.iter().enumerate() {
302            self.roles.get_mut(id).unwrap().rank = rank as i64;
303        }
304
305        db.update_server(
306            &self.id,
307            &PartialServer {
308                roles: Some(self.roles.clone()),
309                ..Default::default()
310            },
311            Vec::new(),
312        )
313        .await?;
314
315        // Publish bulk update event
316        EventV1::ServerRoleRanksUpdate {
317            id: self.id.clone(),
318            ranks: new_order,
319        }
320        .p(self.id.clone())
321        .await;
322
323        Ok(())
324    }
325}
326
327impl Role {
328    /// Into optional struct
329    pub fn into_optional(self) -> PartialRole {
330        PartialRole {
331            id: Some(self.id),
332            name: Some(self.name),
333            permissions: Some(self.permissions),
334            colour: self.colour,
335            hoist: Some(self.hoist),
336            rank: Some(self.rank),
337            icon: self.icon,
338        }
339    }
340
341    /// Create a role
342    pub async fn create(db: &Database, server: &Server, name: String) -> Result<Self> {
343        let role = Role {
344            id: Ulid::new().to_string(),
345            name,
346            // Rank of the new role should be below the lowest role
347            rank: server.roles.len() as i64,
348            colour: None,
349            hoist: false,
350            permissions: Default::default(),
351            icon: None,
352        };
353
354        db.insert_role(&server.id, &role).await?;
355
356        EventV1::ServerRoleUpdate {
357            id: server.id.clone(),
358            role_id: role.id.clone(),
359            data: role.clone().into_optional().into(),
360            clear: vec![],
361        }
362        .p(server.id.clone())
363        .await;
364
365        Ok(role)
366    }
367
368    /// Update server data
369    pub async fn update(
370        &mut self,
371        db: &Database,
372        server_id: &str,
373        partial: PartialRole,
374        remove: Vec<FieldsRole>,
375    ) -> Result<()> {
376        for field in &remove {
377            self.remove_field(field);
378        }
379
380        self.apply_options(partial.clone());
381
382        db.update_role(server_id, &self.id, &partial, remove.clone())
383            .await?;
384
385        EventV1::ServerRoleUpdate {
386            id: server_id.to_string(),
387            role_id: self.id.clone(),
388            data: partial.into(),
389            clear: remove.into_iter().map(Into::into).collect(),
390        }
391        .p(server_id.to_string())
392        .await;
393
394        Ok(())
395    }
396
397    /// Remove field from Role
398    pub fn remove_field(&mut self, field: &FieldsRole) {
399        match field {
400            FieldsRole::Colour => self.colour = None,
401            FieldsRole::Icon => self.icon = None,
402        }
403    }
404
405    /// Generates a PartialRole containing the data which has changed in an update
406    pub fn generate_diff(&self, partial: &PartialRole, remove: &[FieldsRole]) -> PartialRole {
407        let mut before = PartialRole::default();
408
409        generate_diff!(
410            self, before, partial, remove,
411            (
412                name,
413                permissions,
414                (FieldsRole::Colour) colour,
415                hoist,
416                rank,
417                (FieldsRole::Icon) icon,
418            )
419        );
420
421        before
422    }
423
424    /// Delete a role
425    pub async fn delete(&self, db: &Database, server_id: &str) -> Result<()> {
426        EventV1::ServerRoleDelete {
427            id: server_id.to_string(),
428            role_id: self.id.clone(),
429        }
430        .p(server_id.to_string())
431        .await;
432
433        db.delete_role(server_id, &self.id).await
434    }
435}
436
437impl SystemMessageChannels {
438    pub fn into_channel_ids(self) -> HashSet<String> {
439        let mut ids = HashSet::new();
440
441        if let Some(id) = self.user_joined {
442            ids.insert(id);
443        }
444
445        if let Some(id) = self.user_left {
446            ids.insert(id);
447        }
448
449        if let Some(id) = self.user_kicked {
450            ids.insert(id);
451        }
452
453        if let Some(id) = self.user_banned {
454            ids.insert(id);
455        }
456
457        ids
458    }
459}
460
461#[cfg(test)]
462mod tests {
463    use revolt_permissions::{calculate_server_permissions, ChannelPermission};
464
465    use crate::{fixture, util::permissions::DatabasePermissionQuery};
466
467    #[tokio::test]
468    async fn permissions() {
469        database_test!(|db| async move {
470            fixture!(db, "server_with_roles",
471                owner user 0
472                moderator user 1
473                user user 2
474                server server 4);
475
476            let mut query = DatabasePermissionQuery::new(&db, &owner).server(&server);
477            assert!(calculate_server_permissions(&mut query)
478                .await
479                .has_channel_permission(ChannelPermission::GrantAllSafe));
480
481            let mut query = DatabasePermissionQuery::new(&db, &moderator).server(&server);
482            assert!(calculate_server_permissions(&mut query)
483                .await
484                .has_channel_permission(ChannelPermission::BanMembers));
485
486            let mut query = DatabasePermissionQuery::new(&db, &user).server(&server);
487            assert!(!calculate_server_permissions(&mut query)
488                .await
489                .has_channel_permission(ChannelPermission::BanMembers));
490        });
491    }
492}