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::{error, trace};
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, io::{BufWriter, Write}, num::NonZeroU32, ops::{Deref, DerefMut}, path::{Path, PathBuf}, time::Duration
16};
17use tokio::sync::Semaphore;
18
19use crate::{Error, Result};
20
21pub struct ArchiveClientBuilder {
22 client: Client,
23 pre_min_limit: u32,
24 pre_sec_limit: Option<u32>,
25 max_conn_limit: Option<u32>,
26 retry_limit: u32,
27}
28
29impl ArchiveClientBuilder {
30 pub fn new(client: Client, pre_min_limit: u32) -> Self {
31 Self {
32 client,
33 pre_min_limit,
34 pre_sec_limit: None,
35 max_conn_limit: None,
36 retry_limit: 3,
37 }
38 }
39
40 pub fn pre_min_limit(mut self, limit: u32) -> Self {
41 self.pre_min_limit = limit;
42 self
43 }
44
45 pub fn pre_sec_limit(mut self, limit: u32) -> Self {
46 self.pre_sec_limit = Some(limit);
47 self
48 }
49
50 pub fn max_conn_limit(mut self, limit: u32) -> Self {
51 self.max_conn_limit = Some(limit);
52 self
53 }
54
55 pub fn retry_limit(mut self, limit: u32) -> Self {
56 self.retry_limit = limit;
57 self
58 }
59
60 pub fn build(self) -> ArchiveClient {
61 let retry_policy = ExponentialBackoff::builder().build_with_max_retries(self.retry_limit);
62 let client = reqwest_middleware::ClientBuilder::new(self.client)
63 .with(SemaphoreMiddleware::new(
64 self.pre_min_limit,
65 self.pre_sec_limit.or(self.max_conn_limit).unwrap_or(4),
66 self.max_conn_limit.or(self.pre_sec_limit).unwrap_or(4),
67 ))
68 .with(RetryTransientMiddleware::new_with_policy(retry_policy))
69 .build();
70
71 ArchiveClient {
72 inner: client,
73 retry: self.retry_limit,
74 }
75 }
76}
77
78#[derive(Debug, Clone)]
79pub struct ArchiveClient {
80 inner: ClientWithMiddleware,
81 retry: u32,
82}
83
84impl ArchiveClient {
85 pub fn builder(client: Client, pre_min_limit: u32) -> ArchiveClientBuilder {
86 ArchiveClientBuilder::new(client, pre_min_limit)
87 }
88
89 async fn fetch_with_method_without_retry<T: DeserializeOwned>(
90 &self,
91 method: Method,
92 url: impl IntoUrl + Clone,
93 ) -> Result<T> {
94 let request = self.inner.request(method, url);
95 let response = request.send().await?;
96 let response = response.bytes().await?;
97 serde_json::from_slice(&response).map_err(|e| {
98 Error::UnexpectedResponse(e, String::from_utf8(response.to_vec()).unwrap())
99 })
100 }
101
102 pub async fn fetch_with_method<T: DeserializeOwned>(
103 &self,
104 method: Method,
105 url: impl IntoUrl + Clone,
106 ) -> Result<T> {
107 for i in 0..=self.retry {
108 match self
109 .fetch_with_method_without_retry(method.clone(), url.clone())
110 .await
111 {
112 Ok(data) => return Ok(data),
113 Err(e) => {
114 let url = url.clone().into_url()?;
115 let max_retry = self.retry + 1;
116 if i == self.retry {
117 error!("Failed to fetch {url} after {max_retry} attempts: {e}");
118 return Err(e);
119 } else {
120 let retry_count = i + 1;
121 error!(
122 "Attempt {retry_count}/{max_retry} to fetch {url} failed: {e}. Retrying..."
123 );
124 }
125 }
126 }
127 }
128 unreachable!();
129 }
130
131 pub async fn fetch<T: DeserializeOwned>(&self, url: impl IntoUrl + Clone) -> Result<T> {
132 self.fetch_with_method(Method::GET, url).await
133 }
134
135 pub async fn download(
136 &self,
137 with_in: &Path,
138 url: impl IntoUrl + Clone,
139 ) -> Result<PathBuf> {
140 async fn handle(with_in: &Path , request: RequestBuilder) -> Result<PathBuf> {
141 let response = request.send().await?;
142 let mut stream = response.bytes_stream();
143
144 let filename: String = (0..10)
145 .map(|_| fastrand::alphanumeric())
146 .collect();
147 let path = with_in.join(filename);
148 let mut file = File::create(&path)?;
149
150 let mut buffer = BufWriter::new(&mut file);
151 while let Some(bytes) = stream.next().await {
152 let bytes = bytes?;
153 buffer.write_all(&bytes)?;
154 }
155 buffer.flush()?;
156 drop(buffer);
157
158 file.sync_all()?;
159 Ok(path)
160 }
161
162 for i in 0..=self.retry {
163 let request = self.request(Method::GET, url.clone());
164 match handle(&with_in, request).await {
165 Ok(file) => return Ok(file),
166 Err(e) => {
167 let url = url.clone().into_url()?;
168 let max_retry = self.retry + 1;
169 if i == self.retry {
170 error!("Failed to download {url} after {max_retry} attempts: {e}");
171 return Err(e);
172 } else {
173 let retry_count = i + 1;
174 error!(
175 "Attempt {retry_count}/{max_retry} to download {url} failed: {e}. Retrying..."
176 );
177 }
178 }
179 }
180 }
181 unreachable!();
182 }
183}
184
185impl Deref for ArchiveClient {
186 type Target = ClientWithMiddleware;
187
188 fn deref(&self) -> &Self::Target {
189 &self.inner
190 }
191}
192
193impl DerefMut for ArchiveClient {
194 fn deref_mut(&mut self) -> &mut Self::Target {
195 &mut self.inner
196 }
197}
198
199type ArchiveRateLimiter =
200 RateLimiter<NotKeyed, InMemoryState, QuantaClock, NoOpMiddleware<QuantaInstant>>;
201#[derive(Debug)]
202pub struct SemaphoreMiddleware {
203 max_conn_semaphore: Semaphore,
204 pre_sec_limiter: ArchiveRateLimiter,
205 pre_min_limiter: ArchiveRateLimiter,
206}
207
208impl SemaphoreMiddleware {
209 pub fn new(pre_min_limit: u32, pre_sec_limit: u32, max_conn_limit: u32) -> Self {
210 let semaphore = Semaphore::new(max_conn_limit as usize);
211 let min_rate_limiter =
212 RateLimiter::direct(Quota::per_minute(NonZeroU32::new(pre_min_limit).unwrap()));
213 let sec_rate_limiter =
214 RateLimiter::direct(Quota::per_second(NonZeroU32::new(pre_sec_limit).unwrap()));
215 Self {
216 max_conn_semaphore: semaphore,
217 pre_sec_limiter: sec_rate_limiter,
218 pre_min_limiter: min_rate_limiter,
219 }
220 }
221}
222
223#[async_trait::async_trait]
224impl Middleware for SemaphoreMiddleware {
225 async fn handle(
226 &self,
227 req: Request,
228 extensions: &mut http::Extensions,
229 next: Next<'_>,
230 ) -> reqwest_middleware::Result<Response> {
231 let _ = self.max_conn_semaphore.acquire().await.unwrap();
232 self.pre_sec_limiter.until_ready().await;
233 self.pre_min_limiter
234 .until_ready_with_jitter(Jitter::up_to(Duration::from_millis(800)))
235 .await;
236 trace!("Fetching: {}", req.url());
237 next.run(req, extensions).await
238 }
239}