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}