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)]
32 pub struct Notifiers {
33 $($event: WaiterMap<$event_arg>),*
34 }
35
36 impl Notifiers {
37 paste! {
38 $(
39 #[doc = "Waits for the `" $event "` event to be received."]
40 pub async fn [<wait_for_ $event>]<
41 F: Fn(&$event_arg) -> bool + Send + Sync + 'static
42 >(
43 &self,
44 check: F,
45 timeout: Option<Duration>
46 ) -> Result<$event_arg, Error> {
47 self.inner_wait(&self.$event, check, timeout).await
48 }
49
50 #[doc = "Invokes the `" $event "` event."]
51 pub async fn [<invoke_ $event _waiters>](&self, arg: &$event_arg) {
52 self.inner_invoke(&self.$event, arg).await
53 }
54
55 #[doc = "Clears all waiters for the `" $event "` event, this will cause them to return [`Error::BrokenChannel`] early."]
56 pub async fn [<clear_ $event _waiters>](&self) {
57 self.$event.lock().await.clear();
58 }
59 )*
60
61 #[doc = "Clears all waiters for all events, this will cause them to return [`Error::BrokenChannel`] early."]
62 pub async fn clear_all_waiters(&self) {
63 $(
64 self.[<clear_ $event _waiters>]().await;
65 )*
66 }
67 }
68 }
69 }
70}
71
72generate_notifiers! {
73 event: EventV1,
74 authenticated: (),
75 logout: (),
76 pong: Ping,
77 ready: (),
78 message: Message,
79 message_update: (Message, Message, PartialMessage, Vec<FieldsMessage>),
80 message_delete: Message,
81 message_react: (Message, String, String),
82 message_unreact: (Message, String, String),
83 message_remove_reaction: (Message, String, IndexSet<String>),
84 message_append: (Message, Vec<Embed>),
85 user_update: (User, User, PartialUser, Vec<FieldsUser>),
86 bulk_message_delete: (String, Vec<String>, Vec<Message>),
87 channel_create: Channel,
88 channel_update: (Channel, Channel, PartialChannel, Vec<FieldsChannel>),
89 channel_delete: Channel,
90 channel_group_user_join: (Channel, String),
91 channel_group_user_leave: (Channel, String),
92 server_create: (Server, Vec<Channel>, Vec<Emoji>, Vec<ChannelVoiceState>),
93 server_delete: (Server, Vec<Channel>, Vec<Emoji>, Vec<ChannelVoiceState>),
94 server_update: (Server, Server, PartialServer, Vec<FieldsServer>),
95 typing_start: (String, String),
96 typing_stop: (String, String),
97 server_member_join: Member,
98 server_member_leave: (Member, RemovalIntention),
99 server_member_update: (Member, Member, PartialMember, Vec<FieldsMember>),
100 server_role_create: (String, Role),
101 server_role_update: (String, Role, Role, PartialRole, Vec<FieldsRole>),
102 server_role_delete: (String, Role),
103 server_role_ranks_update: (String, Vec<Role>, Vec<Role>),
104 user_voice_state_update: (UserVoiceState, UserVoiceState, PartialUserVoiceState),
105 user_voice_channel_join: (String, UserVoiceState),
106 user_voice_channel_move: (String, String, String, UserVoiceState, UserVoiceState),
107 user_voice_channel_leave: (String, UserVoiceState),
108 emoji_create: Emoji,
109 emoji_delete: Emoji,
110}
111
112impl Notifiers {
113 async fn inner_wait<F: Fn(&M) -> bool + Send + Sync + 'static, M: Clone>(
114 &self,
115 waiters: &WaiterMap<M>,
116 check: F,
117 timeout: Option<Duration>,
118 ) -> Result<M, Error> {
119 let (sender, receiver) = oneshot::channel();
120
121 let random_value = random();
122
123 {
124 let mut lock = waiters.lock().await;
125
126 lock.insert(
127 random_value,
128 Waiter {
129 check: Arc::new(Box::new(check)),
130 oneshot: Arc::new(Mutex::new(Some(sender))),
131 },
132 );
133 }
134
135 let response = if let Some(timeout) = timeout {
136 tokio::time::timeout(timeout, receiver)
137 .await
138 .map(|res| res.map_err(|_| Error::BrokenChannel))
139 .map_err(|_| Error::Timeout)
140 } else {
141 Ok(receiver.await.map_err(|_| Error::BrokenChannel))
142 };
143
144 {
145 let mut lock = waiters.lock().await;
146
147 lock.remove(&random_value);
148 }
149
150 response?
151 }
152
153 async fn inner_invoke<M: Clone + Debug>(&self, waiters: &WaiterMap<M>, value: &M) {
154 let lock = waiters.lock().await.clone();
155
156 for (id, waiter) in lock {
157 if (waiter.check)(value) {
158 if let Some(oneshot) = waiter.oneshot.lock().await.take() {
159 if let Err(e) = oneshot.send(value.clone()) {
160 log::error!("Notifier failed with payload {e:?}")
161 }
162 }
163
164 waiters.lock().await.remove(&id);
165 }
166 }
167 }
168}
169
170impl AsRef<Notifiers> for Notifiers {
171 fn as_ref(&self) -> &Notifiers {
172 self
173 }
174}