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 pub struct Server {
18 #[serde(rename = "_id")]
20 pub id: String,
21 pub owner: String,
23
24 pub name: String,
26 #[serde(skip_serializing_if = "Option::is_none")]
28 pub description: Option<String>,
29
30 pub channels: Vec<String>,
33 #[serde(skip_serializing_if = "Option::is_none")]
35 pub categories: Option<Vec<Category>>,
36 #[serde(skip_serializing_if = "Option::is_none")]
38 pub system_messages: Option<SystemMessageChannels>,
39
40 #[serde(
42 default = "HashMap::<String, Role>::new",
43 skip_serializing_if = "HashMap::<String, Role>::is_empty"
44 )]
45 pub roles: HashMap<String, Role>,
46 pub default_permissions: i64,
48
49 #[serde(skip_serializing_if = "Option::is_none")]
51 pub icon: Option<File>,
52 #[serde(skip_serializing_if = "Option::is_none")]
54 pub banner: Option<File>,
55
56 #[serde(skip_serializing_if = "Option::is_none")]
58 pub flags: Option<i32>,
59
60 #[serde(skip_serializing_if = "crate::if_false", default)]
62 pub nsfw: bool,
63 #[serde(skip_serializing_if = "crate::if_false", default)]
65 pub analytics: bool,
66 #[serde(skip_serializing_if = "crate::if_false", default)]
68 pub discoverable: bool,
69 },
70 "PartialServer"
71);
72
73auto_derived_partial!(
74 pub struct Role {
76 #[serde(rename = "_id")]
78 pub id: String,
79 pub name: String,
81 pub permissions: OverrideField,
83 #[serde(skip_serializing_if = "Option::is_none")]
87 pub colour: Option<String>,
88 #[serde(skip_serializing_if = "crate::if_false", default)]
90 pub hoist: bool,
91 #[serde(default)]
93 pub rank: i64,
94 #[serde(skip_serializing_if = "Option::is_none")]
96 pub icon: Option<File>,
97 },
98 "PartialRole"
99);
100
101auto_derived!(
102 pub struct Category {
104 pub id: String,
106 pub title: String,
108 pub channels: Vec<String>,
110 }
111
112 pub struct SystemMessageChannels {
114 #[serde(skip_serializing_if = "Option::is_none")]
116 pub user_joined: Option<String>,
117 #[serde(skip_serializing_if = "Option::is_none")]
119 pub user_left: Option<String>,
120 #[serde(skip_serializing_if = "Option::is_none")]
122 pub user_kicked: Option<String>,
123 #[serde(skip_serializing_if = "Option::is_none")]
125 pub user_banned: Option<String>,
126 }
127
128 pub enum FieldsServer {
130 Description,
131 Categories,
132 SystemMessages,
133 Icon,
134 Banner,
135 }
136
137 pub enum FieldsRole {
139 Colour,
140 Icon,
141 }
142);
143
144#[allow(clippy::disallowed_methods)]
145impl Server {
146 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 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 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 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 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 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 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 pub async fn set_role_ordering(&mut self, db: &Database, new_order: Vec<String>) -> Result<()> {
302 debug_assert_eq!(self.roles.len(), new_order.len());
304
305 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 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 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 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 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: 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 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 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 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 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}