1use std::borrow::Cow;
4use std::io;
5use std::time::Duration;
6
7use moirai_async::timer::timeout;
8use moirai_tls::TlsConnector;
9
10use crate::codec::{DEFAULT_MAX_RESPONSE_BYTES, read_response, write_request};
11use crate::conn::Conn;
12use crate::pool::IdlePool;
13use crate::redirect::{
14 forwarded_headers, is_redirect, parse_url, redirects_to_get, resolve_redirect,
15};
16use crate::{Origin, Response};
17
18const DEFAULT_MAX_IDLE_PER_ORIGIN: usize = 8;
19const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
20const DEFAULT_MAX_REDIRECTS: usize = 10;
21const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
22
23pub struct HttpClient {
25 tls: TlsConnector,
26 pool: IdlePool<Origin, Conn>,
27 max_idle_per_host: usize,
28 idle_timeout: Duration,
29 max_redirects: usize,
30 request_timeout: Duration,
31 max_response_bytes: usize,
32}
33
34impl Default for HttpClient {
35 fn default() -> Self {
36 Self::new()
37 }
38}
39
40impl HttpClient {
41 #[must_use]
47 pub fn new() -> Self {
48 Self::with_tls(TlsConnector::with_webpki_roots())
49 }
50
51 #[must_use]
53 pub fn with_tls(tls: TlsConnector) -> Self {
54 Self {
55 tls,
56 pool: IdlePool::default(),
57 max_idle_per_host: DEFAULT_MAX_IDLE_PER_ORIGIN,
58 idle_timeout: DEFAULT_IDLE_TIMEOUT,
59 max_redirects: DEFAULT_MAX_REDIRECTS,
60 request_timeout: DEFAULT_REQUEST_TIMEOUT,
61 max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
62 }
63 }
64
65 pub fn set_timeout(&mut self, duration: Duration) {
70 self.request_timeout = duration;
71 }
72
73 pub fn set_max_response_bytes(&mut self, bytes: usize) {
78 self.max_response_bytes = bytes;
79 }
80
81 pub fn set_max_idle_per_host(&mut self, connections: usize) {
85 self.max_idle_per_host = connections;
86 }
87
88 pub fn set_idle_timeout(&mut self, duration: Duration) {
93 self.idle_timeout = duration;
94 }
95
96 pub fn set_max_redirects(&mut self, redirects: usize) {
100 self.max_redirects = redirects;
101 }
102
103 pub async fn get(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
109 self.request("GET", url, headers, None).await
110 }
111
112 pub async fn head(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
118 self.request("HEAD", url, headers, None).await
119 }
120
121 pub async fn request(
138 &self,
139 method: &str,
140 url: &str,
141 headers: &[(&str, &str)],
142 body: Option<&[u8]>,
143 ) -> io::Result<Response> {
144 match timeout(
145 self.request_timeout,
146 self.request_following_redirects(method, url, headers, body),
147 )
148 .await
149 {
150 Ok(result) => result,
151 Err(_) => Err(io::Error::new(
152 io::ErrorKind::TimedOut,
153 "logical HTTP request timed out",
154 )),
155 }
156 }
157
158 async fn request_following_redirects<'a>(
159 &self,
160 method: &'a str,
161 url: &'a str,
162 headers: &'a [(&'a str, &'a str)],
163 body: Option<&'a [u8]>,
164 ) -> io::Result<Response> {
165 let mut current_method = Cow::Borrowed(method);
166 let mut current_url = Cow::Borrowed(url);
167 let mut current_headers: Option<Vec<(&str, &str)>> = None;
168 let mut current_body = body;
169 let mut redirects_followed = 0usize;
170
171 loop {
172 let request_headers = current_headers.as_deref().unwrap_or(headers);
173 let response = self
174 .request_once_url(
175 current_method.as_ref(),
176 current_url.as_ref(),
177 request_headers,
178 current_body,
179 )
180 .await?;
181 if !is_redirect(response.status) {
182 return Ok(response);
183 }
184 let Some(location) = response.header("location") else {
185 return Ok(response);
186 };
187 if redirects_followed >= self.max_redirects {
188 return Err(io::Error::new(
189 io::ErrorKind::InvalidData,
190 format!("redirect limit of {} exceeded", self.max_redirects),
191 ));
192 }
193
194 let (origin, _) = parse_url(current_url.as_ref())?;
195 let next_url = resolve_redirect(current_url.as_ref(), location)?;
196 let (next_origin, _) = parse_url(&next_url)?;
197 let body_dropped = redirects_to_get(response.status, current_method.as_ref());
198 let next_headers =
199 forwarded_headers(request_headers, next_origin != origin, body_dropped);
200 if body_dropped {
201 current_method = Cow::Borrowed("GET");
202 current_body = None;
203 }
204 current_headers = Some(next_headers);
205 current_url = Cow::Owned(next_url);
206 redirects_followed = redirects_followed.checked_add(1).ok_or_else(|| {
207 io::Error::new(io::ErrorKind::InvalidData, "redirect counter overflow")
208 })?;
209 }
210 }
211
212 async fn request_once_url(
213 &self,
214 method: &str,
215 url: &str,
216 headers: &[(&str, &str)],
217 body: Option<&[u8]>,
218 ) -> io::Result<Response> {
219 let (origin, path) = parse_url(url)?;
220 if let Some(connection) = self.pool.take(&origin, self.idle_timeout) {
221 match self
222 .try_once(connection, method, &origin, &path, headers, body)
223 .await
224 {
225 Ok(response) => return Ok(response),
226 Err(pooled_error) if !is_idempotent(method) => return Err(pooled_error),
227 Err(pooled_error) => {
228 let connection = Conn::connect(&origin, &self.tls)
229 .await
230 .map_err(|error| with_retry_context(error, &pooled_error))?;
231 return self
232 .try_once(connection, method, &origin, &path, headers, body)
233 .await
234 .map_err(|error| with_retry_context(error, &pooled_error));
235 }
236 }
237 }
238
239 let connection = Conn::connect(&origin, &self.tls).await?;
240 self.try_once(connection, method, &origin, &path, headers, body)
241 .await
242 }
243
244 async fn try_once(
245 &self,
246 mut connection: Conn,
247 method: &str,
248 origin: &Origin,
249 path: &str,
250 headers: &[(&str, &str)],
251 body: Option<&[u8]>,
252 ) -> io::Result<Response> {
253 let host = origin.host_header();
254 write_request(&mut connection, method, &host, path, headers, body).await?;
255 let response = read_response(
256 &mut connection,
257 method.eq_ignore_ascii_case("HEAD"),
258 self.max_response_bytes,
259 )
260 .await?;
261 if response.keep_alive {
262 self.pool.put(origin, connection, self.max_idle_per_host);
263 }
264 Ok(response)
265 }
266}
267
268fn is_idempotent(method: &str) -> bool {
269 ["GET", "HEAD", "OPTIONS", "TRACE", "PUT", "DELETE"]
270 .iter()
271 .any(|candidate| method.eq_ignore_ascii_case(candidate))
272}
273
274fn with_retry_context(error: io::Error, pooled_error: &io::Error) -> io::Error {
275 io::Error::new(
276 error.kind(),
277 format!("{error} (retry after stale pooled connection failed with: {pooled_error})"),
278 )
279}
280
281#[cfg(test)]
282mod tests;