Skip to main content

stoat/
http.rs

1use std::{
2    hash::{DefaultHasher, Hasher},
3    sync::Arc,
4    time::Duration,
5};
6
7use bytes::Bytes;
8use reqwest::{Client, Method, Request, RequestBuilder, Response};
9use scc::HashMap;
10use serde::{Deserialize, Serialize};
11use stoat_models::v0::{
12    BanListResult, BulkMessageResponse, Channel, CreateVoiceUserResponse, CreateWebhookBody,
13    DataBanCreate, DataCreateRole, DataCreateServerChannel, DataDefaultChannelPermissions,
14    DataEditChannel, DataEditMessage, DataEditRole, DataEditRoleRanks, DataEditServer,
15    DataEditUser, DataJoinCall, DataMemberEdit, DataMessageSend, DataSetRolePermissions,
16    DataSetServerRolePermission, Emoji, FetchServerResponse, FlagResponse, Invite, Member, Message,
17    MutualResponse, NewRoleResponse, OptionsBulkDelete, OptionsFetchServer, OptionsQueryMessages,
18    OptionsServerDelete, OptionsUnreact, Role, Server, ServerBan, User, UserProfile, Webhook,
19};
20use stoat_permissions::DataPermissionsValue;
21use tokio::time::sleep;
22
23use crate::{
24    error::{Error, Result},
25    types::{AutumnResponse, StoatConfig},
26};
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub struct RatelimitEntry {
30    pub remaining: u32,
31    pub reset: u128,
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35enum Service {
36    Api,
37    Autumn,
38}
39
40#[derive(Clone, Debug)]
41pub struct HttpClient {
42    pub base: String,
43    pub api_config: Arc<StoatConfig>,
44    pub token: Option<String>,
45    pub user_id: Option<String>,
46    pub inner: Client,
47    pub ratelimits: Arc<HashMap<u64, RatelimitEntry>>,
48}
49
50impl AsRef<HttpClient> for HttpClient {
51    fn as_ref(&self) -> &HttpClient {
52        self
53    }
54}
55
56impl AsRef<StoatConfig> for HttpClient {
57    fn as_ref(&self) -> &StoatConfig {
58        &self.api_config
59    }
60}
61
62impl AsRef<StoatConfig> for StoatConfig {
63    fn as_ref(&self) -> &StoatConfig {
64        self
65    }
66}
67
68impl HttpClient {
69    pub async fn new(base: String, token: Option<String>, user_id: Option<String>) -> Result<Self> {
70        let client = Client::new();
71        let ratelimits = Arc::new(HashMap::new());
72
73        let api_config = HttpRequest {
74            ratelimits: ratelimits.clone(),
75            service: Service::Api,
76            builder: client.get(&base),
77        }
78        .response()
79        .await?;
80
81        Ok(HttpClient {
82            base,
83            api_config: Arc::new(api_config),
84            token,
85            user_id,
86            inner: client,
87            ratelimits,
88        })
89    }
90
91    pub fn request(&self, method: Method, route: impl AsRef<str>) -> HttpRequest {
92        let mut builder = self
93            .inner
94            .request(method, format!("{}{}", &self.base, route.as_ref()))
95            .header("Accept", "application/json");
96
97        if let Some(token) = &self.token {
98            builder = builder.header("x-bot-token", token);
99        }
100
101        HttpRequest {
102            ratelimits: self.ratelimits.clone(),
103            service: Service::Api,
104            builder,
105        }
106    }
107
108    pub fn autumn_request(&self, method: Method, route: impl AsRef<str>) -> HttpRequest {
109        let mut builder = self
110            .inner
111            .request(
112                method,
113                format!("{}{}", &self.api_config.features.autumn.url, route.as_ref()),
114            )
115            .header("Accept", "application/json");
116
117        if let Some(token) = &self.token {
118            builder = builder.header("x-bot-token", token);
119        }
120
121        HttpRequest {
122            ratelimits: self.ratelimits.clone(),
123            service: Service::Autumn,
124            builder,
125        }
126    }
127
128    pub fn format_file_url(&self, tag: &str, id: &str, filename: Option<&str>) -> String {
129        let autumn_url = &self.api_config.features.autumn.url;
130
131        let mut url = format!("{autumn_url}/{}/{}", tag, id);
132
133        if let Some(filename) = filename {
134            url.push('/');
135            url.push_str(filename);
136        };
137
138        url
139    }
140
141    pub async fn get_root(&self) -> Result<StoatConfig> {
142        self.request(Method::GET, "/").response().await
143    }
144
145    pub async fn send_message(&self, channel_id: &str, data: &DataMessageSend) -> Result<Message> {
146        self.request(Method::POST, format!("/channels/{}/messages", channel_id))
147            .body(data)
148            .response()
149            .await
150    }
151
152    pub async fn fetch_user(&self, user_id: &str) -> Result<User> {
153        self.request(Method::GET, format!("/users/{user_id}"))
154            .response()
155            .await
156    }
157
158    pub async fn fetch_messages<'a>(
159        &self,
160        channel_id: &str,
161        data: &OptionsQueryMessages,
162    ) -> Result<BulkMessageResponse> {
163        self.request(Method::GET, format!("/channels/{}/messages", channel_id))
164            .query(data)
165            .response()
166            .await
167    }
168
169    pub async fn open_dm(&self, user_id: &str) -> Result<Channel> {
170        self.request(Method::GET, format!("/users/{user_id}/dm"))
171            .response()
172            .await
173    }
174
175    pub async fn fetch_member(&self, server_id: &str, user_id: &str) -> Result<Member> {
176        self.request(
177            Method::GET,
178            format!("/servers/{server_id}/members/{user_id}"),
179        )
180        .response()
181        .await
182    }
183
184    pub async fn delete_message(&self, channel_id: &str, message_id: &str) -> Result<()> {
185        self.request(
186            Method::DELETE,
187            format!("/channels/{channel_id}/messages/{message_id}"),
188        )
189        .send()
190        .await
191    }
192
193    pub async fn edit_message(
194        &self,
195        channel_id: &str,
196        message_id: &str,
197        data: &DataEditMessage,
198    ) -> Result<Message> {
199        self.request(
200            Method::PATCH,
201            format!("/channels/{}/messages/{}", channel_id, message_id),
202        )
203        .body(&data)
204        .response()
205        .await
206    }
207
208    pub async fn join_call(
209        &self,
210        channel_id: &str,
211        data: &DataJoinCall,
212    ) -> Result<CreateVoiceUserResponse> {
213        self.request(Method::POST, format!("/channels/{channel_id}/join_call"))
214            .body(data)
215            .response()
216            .await
217    }
218
219    pub async fn fetch_message(&self, channel_id: &str, message_id: &str) -> Result<Message> {
220        self.request(
221            Method::GET,
222            format!("/channels/{channel_id}/messages/{message_id}"),
223        )
224        .response()
225        .await
226    }
227
228    pub async fn upload_file(&self, tag: &str, file: &[u8]) -> Result<AutumnResponse> {
229        self.autumn_request(Method::POST, format!("/{tag}"))
230            .form(&[("file", file)])
231            .response()
232            .await
233    }
234
235    pub async fn delete_channel(&self, channel_id: &str) -> Result<()> {
236        self.request(Method::DELETE, format!("/channels/{channel_id}"))
237            .send()
238            .await
239    }
240
241    pub async fn edit_channel(&self, channel_id: &str, data: &DataEditChannel) -> Result<Channel> {
242        self.request(Method::PATCH, format!("/channels/{channel_id}"))
243            .body(data)
244            .response()
245            .await
246    }
247
248    pub async fn fetch_channel(&self, channel_id: &str) -> Result<Channel> {
249        self.request(Method::GET, format!("/channels/{channel_id}"))
250            .response()
251            .await
252    }
253
254    pub async fn fetch_members(&self, channel_id: &str) -> Result<Vec<User>> {
255        self.request(Method::GET, format!("/channels/{channel_id}/members"))
256            .response()
257            .await
258    }
259
260    pub async fn delete_messages(
261        &self,
262        channel_id: &str,
263        options: &OptionsBulkDelete,
264    ) -> Result<()> {
265        self.request(
266            Method::DELETE,
267            format!("/channels/{channel_id}/messages/bulk"),
268        )
269        .query(options)
270        .send()
271        .await
272    }
273
274    pub async fn clear_reactions(&self, channel_id: &str, message_id: &str) -> Result<()> {
275        self.request(
276            Method::DELETE,
277            format!("/channels/{channel_id}/messages/{message_id}/reactions"),
278        )
279        .send()
280        .await
281    }
282
283    pub async fn pin_message(&self, channel_id: &str, message_id: &str) -> Result<Message> {
284        self.request(
285            Method::POST,
286            format!("/channels/{channel_id}/messages/{message_id}/pin"),
287        )
288        .response()
289        .await
290    }
291
292    pub async fn unpin_message(&self, channel_id: &str, message_id: &str) -> Result<Message> {
293        self.request(
294            Method::DELETE,
295            format!("/channels/{channel_id}/messages/{message_id}/pin"),
296        )
297        .response()
298        .await
299    }
300
301    pub async fn react_message(
302        &self,
303        channel_id: &str,
304        message_id: &str,
305        emoji: &str,
306    ) -> Result<()> {
307        self.request(
308            Method::PUT,
309            format!("/channels/{channel_id}/messages/{message_id}/reactions/{emoji}"),
310        )
311        .send()
312        .await
313    }
314
315    pub async fn unreact_message(
316        &self,
317        channel_id: &str,
318        message_id: &str,
319        emoji: &str,
320        options: &OptionsUnreact,
321    ) -> Result<()> {
322        self.request(
323            Method::DELETE,
324            format!("/channels/{channel_id}/messages/{message_id}/reactions/{emoji}"),
325        )
326        .query(options)
327        .send()
328        .await
329    }
330
331    pub async fn set_default_channel_permissions(
332        &self,
333        channel_id: &str,
334        data: &DataDefaultChannelPermissions,
335    ) -> Result<Channel> {
336        self.request(
337            Method::PUT,
338            format!("/channels/{channel_id}/permissions/default"),
339        )
340        .body(data)
341        .response()
342        .await
343    }
344
345    pub async fn set_role_channel_permissions(
346        &self,
347        channel_id: &str,
348        role_id: &str,
349        data: &DataSetRolePermissions,
350    ) -> Result<Channel> {
351        self.request(
352            Method::PUT,
353            format!("/channels/{channel_id}/permissions/{role_id}"),
354        )
355        .body(data)
356        .response()
357        .await
358    }
359
360    pub async fn create_webhook(
361        &self,
362        channel_id: &str,
363        data: &CreateWebhookBody,
364    ) -> Result<Webhook> {
365        self.request(Method::POST, format!("/channels/{channel_id}/webhooks"))
366            .body(data)
367            .response()
368            .await
369    }
370
371    pub async fn fetch_webhooks(&self, channel_id: &str) -> Result<Vec<Webhook>> {
372        self.request(Method::GET, format!("/channels/{channel_id}/webhooks"))
373            .response()
374            .await
375    }
376
377    pub async fn delete_invite(&self, invite_id: &str) -> Result<()> {
378        self.request(Method::DELETE, format!("/invites/{invite_id}"))
379            .send()
380            .await
381    }
382
383    pub async fn fetch_invite(&self, invite_id: &str) -> Result<Invite> {
384        self.request(Method::GET, format!("/invites/{invite_id}"))
385            .response()
386            .await
387    }
388
389    pub async fn ban_member(
390        &self,
391        server_id: &str,
392        user_id: &str,
393        data: &DataBanCreate,
394    ) -> Result<ServerBan> {
395        self.request(Method::PUT, format!("/servers/{server_id}/bans/{user_id}"))
396            .body(data)
397            .response()
398            .await
399    }
400
401    pub async fn fetch_bans(&self, server_id: &str) -> Result<BanListResult> {
402        self.request(Method::GET, format!("/servers/{server_id}/bans"))
403            .response()
404            .await
405    }
406
407    pub async fn unban_member(&self, server_id: &str, user_id: &str) -> Result<()> {
408        self.request(
409            Method::DELETE,
410            format!("/servers/{server_id}/bans/{user_id}"),
411        )
412        .send()
413        .await
414    }
415
416    pub async fn create_channel(
417        &self,
418        server_id: &str,
419        data: &DataCreateServerChannel,
420    ) -> Result<Channel> {
421        self.request(Method::POST, format!("/servers/{server_id}/channels"))
422            .body(data)
423            .response()
424            .await
425    }
426
427    pub async fn fetch_emojis(&self, server_id: &str) -> Result<Vec<Emoji>> {
428        self.request(Method::GET, format!("/servers/{server_id}/emojis"))
429            .response()
430            .await
431    }
432
433    pub async fn fetch_invites(&self, server_id: &str) -> Result<Vec<Invite>> {
434        self.request(Method::GET, format!("/servers/{server_id}/invites"))
435            .response()
436            .await
437    }
438
439    pub async fn edit_member(
440        &self,
441        server_id: &str,
442        user_id: &str,
443        data: &DataMemberEdit,
444    ) -> Result<Member> {
445        self.request(
446            Method::PATCH,
447            format!("/servers/{server_id}/members/{user_id}"),
448        )
449        .body(data)
450        .response()
451        .await
452    }
453
454    pub async fn kick_member(&self, server_id: &str, user_id: &str) -> Result<()> {
455        self.request(
456            Method::DELETE,
457            format!("/servers/{server_id}/members/{user_id}"),
458        )
459        .send()
460        .await
461    }
462
463    pub async fn set_default_server_permissions(
464        &self,
465        server_id: &str,
466        data: &DataPermissionsValue,
467    ) -> Result<Server> {
468        self.request(
469            Method::PUT,
470            format!("/servers/{server_id}/permissions/default"),
471        )
472        .body(data)
473        .response()
474        .await
475    }
476
477    pub async fn set_role_server_permissions(
478        &self,
479        server_id: &str,
480        role_id: &str,
481        data: &DataSetServerRolePermission,
482    ) -> Result<Server> {
483        self.request(
484            Method::PUT,
485            format!("/servers/{server_id}/permissions/{role_id}"),
486        )
487        .body(data)
488        .response()
489        .await
490    }
491
492    pub async fn create_role(
493        &self,
494        server_id: &str,
495        data: &DataCreateRole,
496    ) -> Result<NewRoleResponse> {
497        self.request(Method::POST, format!("/servers/{server_id}/roles"))
498            .body(data)
499            .response()
500            .await
501    }
502
503    pub async fn delete_role(&self, server_id: &str, role_id: &str) -> Result<()> {
504        self.request(
505            Method::DELETE,
506            format!("/servers/{server_id}/roles/{role_id}"),
507        )
508        .send()
509        .await
510    }
511
512    pub async fn edit_role_positions(
513        &self,
514        server_id: &str,
515        data: &DataEditRoleRanks,
516    ) -> Result<Server> {
517        self.request(Method::PATCH, format!("/servers/{server_id}/roles/ranks"))
518            .body(data)
519            .response()
520            .await
521    }
522
523    pub async fn edit_role(
524        &self,
525        server_id: &str,
526        role_id: &str,
527        data: &DataEditRole,
528    ) -> Result<Role> {
529        self.request(
530            Method::PATCH,
531            format!("/servers/{server_id}/roles/{role_id}"),
532        )
533        .body(data)
534        .response()
535        .await
536    }
537
538    pub async fn fetch_role(&self, server_id: &str, role_id: &str) -> Result<Role> {
539        self.request(Method::GET, format!("/servers/{server_id}/roles/{role_id}"))
540            .response()
541            .await
542    }
543
544    pub async fn delete_server(
545        &self,
546        server_id: &str,
547        options: &OptionsServerDelete,
548    ) -> Result<()> {
549        self.request(Method::DELETE, format!("/servers/{server_id}"))
550            .query(options)
551            .send()
552            .await
553    }
554
555    pub async fn edit_server(&self, server_id: &str, data: &DataEditServer) -> Result<Server> {
556        self.request(Method::PATCH, format!("/servers/{server_id}"))
557            .body(data)
558            .response()
559            .await
560    }
561
562    pub async fn fetch_server(
563        &self,
564        server_id: &str,
565        options: &OptionsFetchServer,
566    ) -> Result<FetchServerResponse> {
567        self.request(Method::GET, format!("/servers/{server_id}"))
568            .query(options)
569            .response()
570            .await
571    }
572
573    pub async fn edit_user(&self, user_id: &str, data: &DataEditUser) -> Result<User> {
574        self.request(Method::PATCH, format!("/users/{user_id}"))
575            .body(data)
576            .response()
577            .await
578    }
579
580    pub async fn fetch_dms(&self) -> Result<Vec<Channel>> {
581        self.request(Method::GET, "/users/dms").response().await
582    }
583
584    pub async fn fetch_user_profile(&self, user_id: &str) -> Result<UserProfile> {
585        self.request(Method::GET, format!("/users/{user_id}/profile"))
586            .response()
587            .await
588    }
589
590    pub async fn fetch_self(&self) -> Result<User> {
591        self.request(Method::GET, "/users/@me").response().await
592    }
593
594    pub async fn fetch_user_flags(&self, user_id: &str) -> Result<FlagResponse> {
595        self.request(Method::GET, format!("/users/{user_id}/flags"))
596            .response()
597            .await
598    }
599
600    pub async fn fetch_user_mutuals(&self, user_id: &str) -> Result<MutualResponse> {
601        self.request(Method::GET, format!("/users/{user_id}/mutual"))
602            .response()
603            .await
604    }
605
606    pub async fn fetch_default_avatar(&self, user_id: &str) -> Result<Bytes> {
607        self.request(Method::GET, format!("/users/{user_id}/default_avatar"))
608            .execute()
609            .await?
610            .bytes()
611            .await
612            .map_err(Into::into)
613    }
614
615    pub async fn fetch_image_preview(&self, tag: &str, id: &str) -> Result<Bytes> {
616        self.autumn_request(Method::GET, format!("{tag}/{id}"))
617            .execute()
618            .await?
619            .bytes()
620            .await
621            .map_err(Into::into)
622    }
623
624    pub async fn fetch_image(&self, tag: &str, id: &str, filename: &str) -> Result<Bytes> {
625        self.autumn_request(Method::GET, format!("{tag}/{id}/{filename}"))
626            .execute()
627            .await?
628            .bytes()
629            .await
630            .map_err(Into::into)
631    }
632}
633
634pub struct HttpRequest {
635    ratelimits: Arc<HashMap<u64, RatelimitEntry>>,
636    service: Service,
637    builder: RequestBuilder,
638}
639
640impl HttpRequest {
641    fn resolve_bucket(service: Service, request: &Request) -> (&str, Option<&str>) {
642        match service {
643            Service::Api => {
644                let mut segments = request.url().path_segments().unwrap();
645
646                let segment = segments.next();
647                let resource = segments.next();
648                let extra = segments.next();
649
650                if let Some(segment) = segment {
651                    let method = request.method();
652
653                    match (segment, resource, method) {
654                        ("users", target, &Method::PATCH) => ("user_edit", target),
655                        ("users", _, _) => {
656                            if let Some("default_avatar") = extra {
657                                return ("default_avatar", None);
658                            }
659
660                            ("users", None)
661                        }
662                        ("bots", _, _) => ("bots", None),
663                        ("channels", Some(id), _) => {
664                            if request.method() == &Method::POST {
665                                if let Some("messages") = extra {
666                                    return ("messaging", Some(id));
667                                }
668                            }
669
670                            ("channels", Some(id))
671                        }
672                        ("servers", Some(id), _) => ("servers", Some(id)),
673                        ("auth", _, _) => {
674                            if request.method() == &Method::DELETE {
675                                ("auth_delete", None)
676                            } else {
677                                ("auth", None)
678                            }
679                        }
680                        ("swagger", _, _) => ("swagger", None),
681                        ("safety", Some("report"), _) => ("safety_report", Some("report")),
682                        ("safety", _, _) => ("safety", None),
683                        _ => ("any", None),
684                    }
685                } else {
686                    ("any", None)
687                }
688            }
689            Service::Autumn => {
690                let path = request.url().path_segments().unwrap().collect::<Vec<_>>();
691
692                match (request.method(), path.as_slice()) {
693                    (&Method::POST, &[tag]) => ("upload", Some(tag)),
694                    _ => ("any", None),
695                }
696            }
697        }
698    }
699
700    pub fn body<I: Serialize>(mut self, body: &I) -> HttpRequest {
701        self.builder = self.builder.json(body);
702
703        self
704    }
705
706    pub fn query<I: Serialize>(mut self, query: &I) -> HttpRequest {
707        self.builder = self.builder.query(query);
708
709        self
710    }
711
712    pub fn form<I: Serialize>(mut self, form: &I) -> HttpRequest {
713        self.builder = self.builder.form(form);
714
715        self
716    }
717
718    pub async fn execute(self) -> Result<Response, Error> {
719        let (client, req) = self.builder.build_split();
720
721        let request = req?;
722
723        let (bucket, resource) = Self::resolve_bucket(self.service, &request);
724        let mut key = DefaultHasher::new();
725        key.write(bucket.as_bytes());
726
727        if let Some(resource) = resource {
728            key.write(resource.as_bytes());
729        };
730
731        let key = key.finish();
732
733        if let Some(entry) = self.ratelimits.get_async(&key).await {
734            if entry.remaining == 0 {
735                let duration = Duration::from_millis(entry.reset as u64);
736
737                log::warn!(
738                    "Ratelimit limit reached: sleeping for {:.3}s",
739                    duration.as_secs_f32()
740                );
741
742                // TODO: switch to queue system to avoid re-hitting the ratelimit after the bucket is reset and too many requests are waiting
743                sleep(duration).await;
744            }
745        }
746
747        let response = client.execute(request).await?;
748
749        let remaining = response
750            .headers()
751            .get("X-RateLimit-Remaining")
752            .unwrap()
753            .to_str()
754            .unwrap()
755            .parse()
756            .unwrap();
757
758        let reset = response
759            .headers()
760            .get("X-RateLimit-Reset-After")
761            .unwrap()
762            .to_str()
763            .unwrap()
764            .parse()
765            .unwrap();
766
767        self.ratelimits
768            .upsert_async(key, RatelimitEntry { remaining, reset })
769            .await;
770
771        if response.status().as_u16() == 429 {
772            let failure = response.json().await?;
773            return Err(Error::RatelimitReached(failure));
774        }
775
776        if response.status().is_client_error() || response.status().is_server_error() {
777            let text = response.json().await?;
778            Err(Error::HttpError(text))
779        } else {
780            Ok(response)
781        }
782    }
783
784    pub async fn response<O: for<'a> Deserialize<'a>>(self) -> Result<O, Error> {
785        self.execute().await?.json().await.map_err(Into::into)
786    }
787
788    pub async fn send(self) -> Result<(), Error> {
789        self.execute().await?;
790
791        Ok(())
792    }
793}