revolt_database/models/channels/ops/
reference.rs1use 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 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 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 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 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 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 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 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 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 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 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 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 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 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}