Skip to main content

stoat/
cache.rs

1use scc::HashMap;
2use std::{
3    collections::VecDeque,
4    sync::{Arc, RwLock},
5};
6use stoat_models::v0::{
7    Channel, ChannelVoiceState, Emoji, EmojiParent, Member, Message, Server, User, UserVoiceState,
8};
9
10use crate::types::{StoatConfig, VoiceNode};
11
12#[derive(Debug, Clone)]
13pub struct GlobalCache {
14    pub api_config: Arc<StoatConfig>,
15
16    pub servers: Arc<HashMap<String, Server>>,
17    pub users: Arc<HashMap<String, User>>,
18    pub members: Arc<HashMap<String, HashMap<String, Member>>>,
19    pub channels: Arc<HashMap<String, Channel>>,
20    pub messages: Arc<RwLock<VecDeque<Message>>>,
21    pub emojis: Arc<HashMap<String, Emoji>>,
22    pub voice_states: Arc<HashMap<String, ChannelVoiceState>>,
23
24    #[cfg(feature = "voice")]
25    pub voice_connections: Arc<HashMap<String, crate::VoiceConnection>>,
26
27    pub current_user_id: Arc<RwLock<Option<String>>>,
28}
29
30impl GlobalCache {
31    pub fn new(api_config: StoatConfig) -> Self {
32        Self {
33            api_config: Arc::new(api_config),
34            servers: Arc::new(HashMap::new()),
35            users: Arc::new(HashMap::new()),
36            members: Arc::new(HashMap::new()),
37            channels: Arc::new(HashMap::new()),
38            messages: Arc::new(RwLock::new(VecDeque::new())),
39            emojis: Arc::new(HashMap::new()),
40            voice_states: Arc::new(HashMap::new()),
41
42            #[cfg(feature = "voice")]
43            voice_connections: Arc::new(HashMap::new()),
44
45            current_user_id: Arc::new(RwLock::new(None)),
46        }
47    }
48
49    pub async fn cleanup(&self) {
50        self.servers.clear_async().await;
51        self.users.clear_async().await;
52        self.members.clear_async().await;
53        self.channels.clear_async().await;
54        self.messages.write().unwrap().clear();
55        self.emojis.clear_async().await;
56        self.voice_states.clear_async().await;
57
58        #[cfg(feature = "voice")]
59        {
60            use futures::{FutureExt, future::join_all};
61
62            let voice_connections = self.voice_connections.clone();
63            let mut iter = voice_connections.begin_async().await;
64            let mut conns = Vec::new();
65
66            while let Some(entry) = iter {
67                let ((_, conn), next) = entry.remove_and_async().await;
68                iter = next;
69                conns.push(conn);
70            }
71
72            join_all(conns.iter().map(|c| c.disconnect().boxed())).await;
73        }
74    }
75
76    pub fn autumn_url(&self) -> &str {
77        &self.api_config.features.autumn.url
78    }
79
80    pub fn livekit_nodes(&self) -> &[VoiceNode] {
81        &self.api_config.features.livekit.nodes
82    }
83
84    pub fn get_server(&self, server_id: &str) -> Option<Server> {
85        self.servers.get_sync(server_id).map(|r| r.get().clone())
86    }
87
88    pub fn insert_server(&self, server: Server) {
89        self.servers.upsert_sync(server.id.clone(), server);
90    }
91
92    pub fn update_server_with<R>(
93        &self,
94        server_id: &str,
95        f: impl FnOnce(&mut Server) -> R,
96    ) -> Option<R> {
97        self.servers.get_sync(server_id).map(|mut r| f(r.get_mut()))
98    }
99
100    pub fn remove_server(&self, server_id: &str) -> Option<Server> {
101        self.servers
102            .remove_sync(server_id)
103            .map(|(_, server)| server)
104    }
105
106    pub fn get_user(&self, user_id: &str) -> Option<User> {
107        self.users.get_sync(user_id).map(|r| r.get().clone())
108    }
109
110    pub fn insert_user(&self, user: User) {
111        self.users.upsert_sync(user.id.clone(), user);
112    }
113
114    pub fn update_user_with<R>(&self, user_id: &str, f: impl FnOnce(&mut User) -> R) -> Option<R> {
115        self.users.get_sync(user_id).map(|mut r| f(r.get_mut()))
116    }
117
118    pub fn remove_user(&self, user_id: &str) -> Option<User> {
119        self.users.remove_sync(user_id).map(|(_, server)| server)
120    }
121
122    pub fn get_member(&self, server_id: &str, user_id: &str) -> Option<Member> {
123        self.members
124            .get_sync(server_id)
125            .and_then(|members| members.get_sync(user_id).map(|r| r.get().clone()))
126    }
127
128    pub fn insert_member(&self, member: Member) {
129        self.members
130            .entry_sync(member.id.server.clone())
131            .or_default()
132            .get_mut()
133            .upsert_sync(member.id.user.clone(), member);
134    }
135
136    pub fn update_member_with<R>(
137        &self,
138        server_id: &str,
139        user_id: &str,
140        f: impl FnOnce(&mut Member) -> R,
141    ) -> Option<R> {
142        self.members
143            .get_sync(server_id)
144            .and_then(|members| members.get_sync(user_id).map(|mut r| f(r.get_mut())))
145    }
146
147    pub fn remove_member(&self, server_id: &str, user_id: &str) -> Option<Member> {
148        self.members
149            .get_sync(server_id)
150            .and_then(|members| members.remove_sync(user_id).map(|(_, member)| member))
151    }
152
153    pub fn get_channel(&self, channel_id: &str) -> Option<Channel> {
154        self.channels.get_sync(channel_id).map(|r| r.get().clone())
155    }
156
157    pub fn insert_channel(&self, channel: Channel) {
158        self.channels.upsert_sync(channel.id().to_string(), channel);
159    }
160
161    pub fn update_channel_with<R>(
162        &self,
163        channel_id: &str,
164        f: impl FnOnce(&mut Channel) -> R,
165    ) -> Option<R> {
166        self.channels
167            .get_sync(channel_id)
168            .map(|mut r| f(r.get_mut()))
169    }
170
171    pub fn remove_channel(&self, channel_id: &str) -> Option<Channel> {
172        self.channels
173            .remove_sync(channel_id)
174            .map(|(_, channel)| channel)
175    }
176
177    pub fn get_message(&self, message_id: &str) -> Option<Message> {
178        self.messages
179            .read()
180            .unwrap()
181            .iter()
182            .find(|msg| &msg.id == message_id)
183            .cloned()
184    }
185
186    pub fn insert_message(&self, message: Message) {
187        let mut messages = self.messages.write().unwrap();
188
189        messages.push_front(message);
190
191        if messages.len() > 1000 {
192            messages.pop_back();
193        }
194    }
195
196    pub fn update_message_with<R>(
197        &self,
198        message_id: &str,
199        f: impl FnOnce(&mut Message) -> R,
200    ) -> Option<R> {
201        self.messages
202            .write()
203            .unwrap()
204            .iter_mut()
205            .find(|msg| &msg.id == message_id)
206            .map(f)
207    }
208
209    pub fn remove_message(&self, message_id: &str) -> Option<Message> {
210        let mut messages = self.messages.write().unwrap();
211
212        if let Some((idx, _)) = messages
213            .iter()
214            .enumerate()
215            .find(|(_, msg)| &msg.id == message_id)
216        {
217            messages.remove(idx)
218        } else {
219            None
220        }
221    }
222
223    pub fn remove_messages(&self, message_ids: &[String]) -> Vec<Message> {
224        let mut channel_messages = self.messages.write().unwrap();
225
226        let mut i = 0;
227        let end = channel_messages.len();
228
229        let mut messages = Vec::new();
230
231        while i < channel_messages.len() - end {
232            if message_ids.contains(&channel_messages[i].id) {
233                messages.push(channel_messages.remove(i).unwrap());
234            } else {
235                i += 1;
236            };
237        }
238
239        messages
240    }
241
242    pub fn get_current_user(&self) -> Option<User> {
243        self.users
244            .get_sync(&self.get_current_user_id()?)
245            .map(|r| r.get().clone())
246    }
247
248    pub fn insert_voice_state(&self, voice_state: ChannelVoiceState) {
249        self.voice_states
250            .upsert_sync(voice_state.id.clone(), voice_state);
251    }
252
253    pub fn remove_voice_state(&self, channel_id: &str) -> Option<ChannelVoiceState> {
254        self.voice_states
255            .remove_sync(channel_id)
256            .map(|(_, voice_state)| voice_state)
257    }
258
259    pub fn get_voice_state(&self, channel_id: &str) -> Option<ChannelVoiceState> {
260        self.voice_states
261            .get_sync(channel_id)
262            .map(|r| r.get().clone())
263    }
264
265    pub fn insert_voice_state_partipant(&self, channel_id: &str, user_voice_state: UserVoiceState) {
266        let mut channel_voice_state = self
267            .voice_states
268            .entry_sync(channel_id.to_string())
269            .or_insert_with(|| ChannelVoiceState {
270                id: channel_id.to_string(),
271                participants: Vec::new(),
272            });
273
274        channel_voice_state
275            .participants
276            .retain(|state| state.id != user_voice_state.id);
277
278        channel_voice_state.participants.push(user_voice_state);
279    }
280
281    pub fn remove_voice_state_partipant(
282        &self,
283        channel_id: &str,
284        user_id: &str,
285    ) -> Option<UserVoiceState> {
286        if let Some(mut channel_voice_state) = self.voice_states.get_sync(channel_id) {
287            if let Some((i, _)) = channel_voice_state
288                .participants
289                .iter()
290                .enumerate()
291                .find(|(_, state)| &state.id == user_id)
292            {
293                Some(channel_voice_state.participants.remove(i))
294            } else {
295                None
296            }
297        } else {
298            None
299        }
300    }
301
302    pub fn update_voice_state_partipant_with<R>(
303        &self,
304        channel_id: &str,
305        user_id: &str,
306        f: impl FnOnce(&mut UserVoiceState) -> R,
307    ) -> Option<R> {
308        if let Some(mut channel_voice_state) = self.voice_states.get_sync(channel_id) {
309            channel_voice_state
310                .participants
311                .iter_mut()
312                .find(|p| p.id == user_id)
313                .map(f)
314        } else {
315            None
316        }
317    }
318
319    #[cfg(feature = "voice")]
320    pub fn insert_voice_connection(&self, connection: crate::VoiceConnection) {
321        self.voice_connections
322            .upsert_sync(connection.channel_id(), connection);
323    }
324
325    #[cfg(feature = "voice")]
326    pub fn remove_voice_connection(&self, channel_id: &str) -> Option<crate::VoiceConnection> {
327        self.voice_connections
328            .remove_sync(channel_id)
329            .map(|(_, voice_connection)| voice_connection)
330    }
331
332    pub fn insert_emoji(&self, emoji: Emoji) {
333        self.emojis.upsert_sync(emoji.id.clone(), emoji);
334    }
335
336    pub fn get_emoji(&self, emoji_id: &str) -> Option<Emoji> {
337        self.emojis.get_sync(emoji_id).map(|r| r.get().clone())
338    }
339
340    pub fn remove_emoji(&self, emoji_id: &str) -> Option<Emoji> {
341        self.emojis.remove_sync(emoji_id).map(|(_, emoji)| emoji)
342    }
343
344    pub fn remove_server_emojis(&self, server_id: &str) -> Vec<Emoji> {
345        let parent = EmojiParent::Server {
346            id: server_id.to_string(),
347        };
348
349        let mut emojis = Vec::new();
350
351        // Workaround for no extract_if alternative
352        self.emojis.retain_sync(|_, emoji| {
353            if &emoji.parent == &parent {
354                emojis.push(emoji.clone());
355
356                true
357            } else {
358                false
359            }
360        });
361
362        emojis
363    }
364
365    pub fn set_current_user_id(&self, user_id: String) {
366        *self.current_user_id.write().unwrap() = Some(user_id);
367    }
368
369    pub fn get_current_user_id(&self) -> Option<String> {
370        self.current_user_id
371            .read()
372            .unwrap()
373            .as_ref()
374            .map(|v| v.clone())
375    }
376}
377
378impl AsRef<GlobalCache> for GlobalCache {
379    fn as_ref(&self) -> &GlobalCache {
380        self
381    }
382}