Skip to main content

stoat/
notifiers.rs

1use std::{collections::HashMap, fmt::Debug, sync::Arc, time::Duration};
2
3use futures::lock::Mutex;
4use indexmap::IndexSet;
5use paste::paste;
6use rand::random;
7use stoat_database::events::client::{EventV1, Ping};
8use stoat_models::v0::{
9    Channel, ChannelVoiceState, Embed, Emoji, FieldsChannel, FieldsMember, FieldsMessage,
10    FieldsRole, FieldsServer, FieldsUser, Member, Message, PartialChannel, PartialMember,
11    PartialMessage, PartialRole, PartialServer, PartialUser, PartialUserVoiceState,
12    RemovalIntention, Role, Server, User, UserVoiceState,
13};
14use tokio::sync::oneshot;
15
16use crate::Error;
17
18#[derive(Clone)]
19struct Waiter<Arg> {
20    check: Arc<Box<dyn Fn(&Arg) -> bool + Send + Sync + 'static>>,
21    oneshot: Arc<Mutex<Option<oneshot::Sender<Arg>>>>,
22}
23
24type WaiterMap<M> = Arc<Mutex<HashMap<usize, Waiter<M>>>>;
25
26macro_rules! generate_notifiers {
27    ($($event: ident: $event_arg: ty),* $(,)?) => {
28        #[derive(Default, Debug, Clone)]
29        pub struct Notifiers {
30            $($event: WaiterMap<$event_arg>),*
31        }
32
33        impl Notifiers {
34            paste! {
35                $(
36                    pub async fn [<wait_for_ $event>]<
37                        F: Fn(&$event_arg) -> bool + Send + Sync + 'static
38                    >(
39                        &self,
40                        check: F,
41                        timeout: Option<Duration>
42                    ) -> Result<$event_arg, Error> {
43                        self.inner_wait(&self.$event, check, timeout).await
44                    }
45
46                    pub async fn [<invoke_ $event _waiters>](&self, arg: &$event_arg) {
47                        self.inner_invoke(&self.$event, arg).await
48                    }
49
50                    pub async fn [<clear_ $event _waiters>](&self) {
51                        self.$event.lock().await.clear();
52                    }
53                )*
54
55                pub async fn clear_all_waiters(&self) {
56                    $(
57                        self.[<clear_ $event _waiters>]().await;
58                    )*
59                }
60            }
61        }
62    }
63}
64
65generate_notifiers! {
66    event: EventV1,
67    authenticated: (),
68    logout: (),
69    pong: Ping,
70    ready: (),
71    message: Message,
72    message_update: (Message, Message, PartialMessage, Vec<FieldsMessage>),
73    message_delete: Message,
74    message_react: (Message, String, String),
75    message_unreact: (Message, String, String),
76    message_remove_reaction: (Message, String, IndexSet<String>),
77    message_append: (Message, Vec<Embed>),
78    user_update: (User, User, PartialUser, Vec<FieldsUser>),
79    bulk_message_delete: (String, Vec<String>, Vec<Message>),
80    channel_create: Channel,
81    channel_update: (Channel, Channel, PartialChannel, Vec<FieldsChannel>),
82    channel_delete: Channel,
83    channel_group_user_join: (Channel, String),
84    channel_group_user_leave: (Channel, String),
85    server_create: (Server, Vec<Channel>, Vec<Emoji>, Vec<ChannelVoiceState>),
86    server_delete: (Server, Vec<Channel>, Vec<Emoji>, Vec<ChannelVoiceState>),
87    server_update: (Server, Server, PartialServer, Vec<FieldsServer>),
88    typing_start: (String, String),
89    typing_stop: (String, String),
90    server_member_join: Member,
91    server_member_leave: (Member, RemovalIntention),
92    server_member_update: (Member, Member, PartialMember, Vec<FieldsMember>),
93    server_role_create: (String, Role),
94    server_role_update: (String, Role, Role, PartialRole, Vec<FieldsRole>),
95    server_role_delete: (String, Role),
96    server_role_ranks_update: (String, Vec<Role>, Vec<Role>),
97    user_voice_state_update: (UserVoiceState, UserVoiceState, PartialUserVoiceState),
98    user_voice_channel_join: (String, UserVoiceState),
99    user_voice_channel_move: (String, String, String, UserVoiceState, UserVoiceState),
100    user_voice_channel_leave: (String, UserVoiceState),
101    emoji_create: Emoji,
102    emoji_delete: Emoji,
103}
104
105impl Notifiers {
106    async fn inner_wait<F: Fn(&M) -> bool + Send + Sync + 'static, M: Clone>(
107        &self,
108        waiters: &WaiterMap<M>,
109        check: F,
110        timeout: Option<Duration>,
111    ) -> Result<M, Error> {
112        let (sender, receiver) = oneshot::channel();
113
114        let random_value = random();
115
116        {
117            let mut lock = waiters.lock().await;
118
119            lock.insert(
120                random_value,
121                Waiter {
122                    check: Arc::new(Box::new(check)),
123                    oneshot: Arc::new(Mutex::new(Some(sender))),
124                },
125            );
126        }
127
128        let response = if let Some(timeout) = timeout {
129            tokio::time::timeout(timeout, receiver)
130                .await
131                .map(|res| res.map_err(|_| Error::BrokenChannel))
132                .map_err(|_| Error::Timeout)
133        } else {
134            Ok(receiver.await.map_err(|_| Error::BrokenChannel))
135        };
136
137        {
138            let mut lock = waiters.lock().await;
139
140            lock.remove(&random_value);
141        }
142
143        response?
144    }
145
146    async fn inner_invoke<M: Clone + Debug>(&self, waiters: &WaiterMap<M>, value: &M) {
147        let lock = waiters.lock().await.clone();
148
149        for (id, waiter) in lock {
150            if (waiter.check)(value) {
151                if let Some(oneshot) = waiter.oneshot.lock().await.take() {
152                    if let Err(e) = oneshot.send(value.clone()) {
153                        log::error!("Notifier failed with payload {e:?}")
154                    }
155                }
156
157                waiters.lock().await.remove(&id);
158            }
159        }
160    }
161}
162
163impl AsRef<Notifiers> for Notifiers {
164    fn as_ref(&self) -> &Notifiers {
165        self
166    }
167}