Skip to main content

stoat/commands/
context.rs

1use std::{borrow::Cow, fmt::Debug, ops::Deref, sync::Arc};
2
3use state::TypeMap;
4use stoat_models::v0::{Channel, Member, Message, Server, User};
5use stoat_permissions::{
6    PermissionValue, calculate_channel_permissions, calculate_server_permissions,
7};
8
9use crate::{
10    Context as MessageContext, Error, GlobalCache, HttpClient, UserExt,
11    builders::SendMessageBuilder,
12    commands::{Command, HelpCommand, Words, handler::Commands},
13    context::Events,
14    notifiers::Notifiers,
15    permissions::user_permissions_query,
16};
17
18type SendSyncMap = TypeMap![Send + Sync];
19
20#[derive(Debug, Clone)]
21pub struct Context<
22    E: From<Error> + Clone + Debug + Send + Sync + 'static,
23    S: Debug + Clone + Send + Sync + 'static,
24> {
25    pub(crate) inner: MessageContext,
26    pub prefix: Option<String>,
27    pub command: Option<Command<E, S>>,
28    pub message: Message,
29    pub state: S,
30    pub words: Words,
31    pub commands: Commands<E, S>,
32    pub help_command: Arc<dyn HelpCommand<E, S>>,
33    pub(crate) local_state: Arc<SendSyncMap>,
34}
35
36impl<
37    E: From<Error> + Clone + Debug + Send + Sync + 'static,
38    S: Debug + Clone + Send + Sync + 'static,
39> Context<E, S>
40{
41    pub fn local_cache<F: FnOnce() -> T, T: Send + Sync + 'static>(&self, f: F) -> &T {
42        self.local_state.try_get().unwrap_or_else(|| {
43            self.local_state.set(f());
44            self.local_state.get()
45        })
46    }
47
48    pub async fn local_cache_async<Fut: Future<Output = T>, T: Send + Sync + 'static>(
49        &self,
50        fut: Fut,
51    ) -> &T {
52        match self.local_state.try_get() {
53            Some(s) => s,
54            None => {
55                self.local_state.set(fut.await);
56                self.local_state.get()
57            }
58        }
59    }
60
61    pub fn get_current_channel(&self) -> Result<Channel, Error> {
62        self.local_cache(|| {
63            struct CurrentChannel(Result<Channel, Error>);
64
65            CurrentChannel(
66                self.cache
67                    .get_channel(&self.message.channel)
68                    .ok_or(Error::InternalError),
69            )
70        })
71        .0
72        .clone()
73    }
74
75    pub fn get_current_server(&self) -> Result<Server, Error> {
76        self.local_cache(|| {
77            struct CurrentServer(Result<Server, Error>);
78
79            CurrentServer(
80                if let Ok(Channel::TextChannel { server, .. }) = self.get_current_channel() {
81                    self.cache.get_server(&server).ok_or(Error::InternalError)
82                } else {
83                    Err(Error::NotInServer)
84                },
85            )
86        })
87        .0
88        .clone()
89    }
90
91    pub async fn get_user(&self) -> Result<User, Error> {
92        self.local_cache_async({
93            struct CurrentUser(Result<User, Error>);
94
95            async move {
96                CurrentUser(if let Some(user) = self.message.user.as_ref() {
97                    Ok(user.clone())
98                } else if let Some(user) = self.cache.get_user(&self.message.author) {
99                    Ok(user.clone())
100                } else {
101                    self.http.fetch_user(&self.message.author).await
102                })
103            }
104        })
105        .await
106        .0
107        .clone()
108    }
109
110    pub async fn get_member(&self) -> Result<Member, Error> {
111        self.local_cache_async({
112            struct CurrentMember(Result<Member, Error>);
113
114            async move {
115                CurrentMember(if let Some(member) = self.message.member.as_ref() {
116                    Ok(member.clone())
117                } else {
118                    match self.get_current_server() {
119                        Ok(server) => {
120                            if let Some(member) =
121                                self.cache.get_member(&server.id, &self.message.author)
122                            {
123                                Ok(member.clone())
124                            } else {
125                                self.http
126                                    .fetch_member(&server.id, &self.message.author)
127                                    .await
128                            }
129                        }
130                        Err(e) => Err(e),
131                    }
132                })
133            }
134        })
135        .await
136        .0
137        .clone()
138    }
139
140    pub async fn get_author_channel_permissions(&self) -> PermissionValue {
141        self.local_cache_async(async {
142            struct ChannelPermissions(PermissionValue);
143
144            let Ok(user) = self.get_user().await else {
145                return ChannelPermissions(0u64.into());
146            };
147            let member = self.get_member().await;
148            let Ok(channel) = self.get_current_channel() else {
149                return ChannelPermissions(0u64.into());
150            };
151            let server = self.get_current_server();
152
153            let mut query =
154                user_permissions_query(self.cache.clone(), self.http.clone(), Cow::Owned(user))
155                    .channel(Cow::Owned(channel));
156
157            if let Ok(server) = server {
158                query = query.server(Cow::Owned(server))
159            };
160
161            if let Ok(member) = member {
162                query = query.member(Cow::Owned(member))
163            };
164
165            ChannelPermissions(calculate_channel_permissions(&mut query).await)
166        })
167        .await
168        .0
169    }
170
171    pub async fn get_author_server_permissions(&self) -> PermissionValue {
172        self.local_cache_async(async {
173            struct ServerPermissions(PermissionValue);
174
175            let Ok(user) = self.get_user().await else {
176                return ServerPermissions(0u64.into());
177            };
178            let member = self.get_member().await;
179            let server = self.get_current_server();
180
181            let mut query =
182                user_permissions_query(self.cache.clone(), self.http.clone(), Cow::Owned(user));
183
184            if let Ok(server) = server {
185                query = query.server(Cow::Owned(server))
186            };
187
188            if let Ok(member) = member {
189                query = query.member(Cow::Owned(member))
190            };
191
192            ServerPermissions(calculate_server_permissions(&mut query).await)
193        })
194        .await
195        .0
196    }
197
198    pub fn clean_prefix(&self) -> String {
199        let Some(ref prefix) = self.prefix else {
200            return String::new();
201        };
202
203        let user = self.cache.get_current_user().unwrap();
204
205        prefix.replace(
206            &format!("<@{}>", &user.id),
207            &user.name().replace("\\", "\\\\"),
208        )
209    }
210
211    pub fn send(&self) -> SendMessageBuilder {
212        SendMessageBuilder::new(self.http.clone(), self.message.channel.clone())
213    }
214
215    pub fn reply(&self, mention: bool) -> SendMessageBuilder {
216        let mut builder = SendMessageBuilder::new(self.http.clone(), self.message.channel.clone());
217        builder.reply(self.message.id.clone(), mention);
218        builder
219    }
220}
221
222impl<
223    E: From<Error> + Clone + Debug + Send + Sync + 'static,
224    S: Debug + Clone + Send + Sync + 'static,
225> Deref for Context<E, S>
226{
227    type Target = MessageContext;
228
229    fn deref(&self) -> &Self::Target {
230        &self.inner
231    }
232}
233
234impl<
235    E: From<Error> + Clone + Debug + Send + Sync + 'static,
236    S: Debug + Clone + Send + Sync + 'static,
237> AsRef<GlobalCache> for Context<E, S>
238{
239    fn as_ref(&self) -> &GlobalCache {
240        &self.cache
241    }
242}
243
244impl<
245    E: From<Error> + Clone + Debug + Send + Sync + 'static,
246    S: Debug + Clone + Send + Sync + 'static,
247> AsRef<HttpClient> for Context<E, S>
248{
249    fn as_ref(&self) -> &HttpClient {
250        &self.http
251    }
252}
253
254impl<
255    E: From<Error> + Clone + Debug + Send + Sync + 'static,
256    S: Debug + Clone + Send + Sync + 'static,
257> AsRef<Notifiers> for Context<E, S>
258{
259    fn as_ref(&self) -> &Notifiers {
260        &self.notifiers
261    }
262}
263
264impl<
265    E: From<Error> + Clone + Debug + Send + Sync + 'static,
266    S: Debug + Clone + Send + Sync + 'static,
267> AsRef<Events> for Context<E, S>
268{
269    fn as_ref(&self) -> &Events {
270        &self.events
271    }
272}