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