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