Skip to main content

revolt_database/models/servers/
model.rs

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