post_archiver_utils/
request.rs

1use futures::StreamExt;
2use governor::{
3    Jitter, Quota, RateLimiter,
4    clock::{QuantaClock, QuantaInstant},
5    middleware::NoOpMiddleware,
6    state::{InMemoryState, NotKeyed},
7};
8use http::Method;
9use log::info;
10use reqwest::{Client, IntoUrl, Request, Response};
11use reqwest_middleware::{ClientWithMiddleware, Middleware, Next, RequestBuilder};
12use reqwest_retry::{RetryTransientMiddleware, policies::ExponentialBackoff};
13use serde::de::DeserializeOwned;
14use std::{
15    fs::File,
16    io::{BufWriter, Write},
17    num::NonZeroU32,
18    ops::{Deref, DerefMut},
19    sync::Arc,
20    time::Duration,
21};
22use tokio::sync::Semaphore;
23
24use crate::Result;
25
26const RETRY_LIMIT: u32 = 3;
27
28#[derive(Debug, Clone)]
29pub struct ArchiveClient(ClientWithMiddleware);
30
31impl ArchiveClient {
32    pub fn new(client: Client, limit: usize) -> Self {
33        let retry_policy = ExponentialBackoff::builder().build_with_max_retries(RETRY_LIMIT);
34        let client = reqwest_middleware::ClientBuilder::new(client)
35            .with(SemaphoreMiddleware::new(limit))
36            .with(RetryTransientMiddleware::new_with_policy(retry_policy))
37            .build();
38
39        Self(client)
40    }
41
42    pub async fn fetch_with_method<T: DeserializeOwned>(
43        &self,
44        method: Method,
45        url: impl IntoUrl,
46    ) -> Result<T> {
47        let request = self.0.request(method, url);
48        let response = request.send().await?;
49        let response = response.bytes().await?;
50        serde_json::from_slice(&response).map_err(Into::into)
51    }
52
53    pub async fn fetch<T: DeserializeOwned>(&self, url: impl IntoUrl) -> Result<T> {
54        self.fetch_with_method(Method::GET, url).await
55    }
56
57    pub async fn download_with_method(
58        &self,
59        method: Method,
60        url: impl IntoUrl + Clone,
61        file: &mut File,
62    ) -> Result<()> {
63        async fn handle(request: RequestBuilder, file: &mut File) -> Result<()> {
64            file.set_len(0)?;
65
66            let response = request.send().await?;
67            let mut stream = response.bytes_stream();
68
69            let mut buffer = BufWriter::new(file);
70            while let Some(bytes) = stream.next().await {
71                let bytes = bytes?;
72                buffer.write_all(&bytes)?;
73            }
74            buffer.flush()?;
75            Ok(())
76        }
77
78        let mut err = Ok(());
79        for _ in 0..=RETRY_LIMIT {
80            let request = self.0.request(method.clone(), url.clone());
81            match handle(request, file).await {
82                Ok(_) => return Ok(()),
83                Err(e) => err = Err(e),
84            }
85        }
86        err
87    }
88
89    pub async fn download(&self, url: impl IntoUrl + Clone, file: &mut File) -> Result<()> {
90        self.download_with_method(Method::GET, url, file).await
91    }
92}
93
94impl Deref for ArchiveClient {
95    type Target = ClientWithMiddleware;
96
97    fn deref(&self) -> &Self::Target {
98        &self.0
99    }
100}
101
102impl DerefMut for ArchiveClient {
103    fn deref_mut(&mut self) -> &mut Self::Target {
104        &mut self.0
105    }
106}
107
108type ArchiveRateLimiter =
109    RateLimiter<NotKeyed, InMemoryState, QuantaClock, NoOpMiddleware<QuantaInstant>>;
110#[derive(Debug, Clone)]
111pub struct SemaphoreMiddleware(Arc<(Semaphore, ArchiveRateLimiter)>);
112
113impl SemaphoreMiddleware {
114    pub fn new(limit: usize) -> Self {
115        let semaphore = Semaphore::new(5);
116        let rate_limiter =
117            RateLimiter::direct(Quota::per_minute(NonZeroU32::new(limit as u32).unwrap()));
118        Self(Arc::new((semaphore, rate_limiter)))
119    }
120}
121
122#[async_trait::async_trait]
123impl Middleware for SemaphoreMiddleware {
124    async fn handle(
125        &self,
126        req: Request,
127        extensions: &mut http::Extensions,
128        next: Next<'_>,
129    ) -> reqwest_middleware::Result<Response> {
130        let (semaphore, rate_limiter) = self.0.as_ref();
131        let _ = semaphore.acquire().await.unwrap();
132        rate_limiter
133            .until_ready_with_jitter(Jitter::up_to(Duration::from_millis(800)))
134            .await;
135        next.run(req, extensions).await
136    }
137}