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 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}