1#![forbid(unsafe_code)]
11#![deny(missing_docs)]
12
13pub mod codec;
14pub mod conn;
15
16pub use codec::Response;
17pub use conn::Origin;
18
19use std::collections::HashMap;
20use std::io;
21use std::str::FromStr;
22use std::sync::Mutex;
23use std::time::Duration;
24
25use codec::{read_response, write_request, DEFAULT_MAX_RESPONSE_BYTES};
26use conn::Conn;
27use moirai_async::timer::timeout;
28use moirai_tls::TlsConnector;
29
30pub struct HttpClient {
32 tls: TlsConnector,
33 pool: Mutex<HashMap<Origin, Vec<Conn>>>,
34 max_idle_per_host: usize,
35 request_timeout: Duration,
36 max_response_bytes: usize,
37}
38
39impl Default for HttpClient {
40 fn default() -> Self {
41 Self::new()
42 }
43}
44
45impl HttpClient {
46 #[must_use]
48 pub fn new() -> Self {
49 Self::with_tls(TlsConnector::with_webpki_roots())
50 }
51
52 #[must_use]
54 pub fn with_tls(tls: TlsConnector) -> Self {
55 Self {
56 tls,
57 pool: Mutex::new(HashMap::new()),
58 max_idle_per_host: 8,
59 request_timeout: Duration::from_secs(30),
60 max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
61 }
62 }
63
64 pub fn set_timeout(&mut self, dur: Duration) {
66 self.request_timeout = dur;
67 }
68
69 pub fn set_max_response_bytes(&mut self, n: usize) {
73 self.max_response_bytes = n;
74 }
75
76 pub fn set_max_idle_per_host(&mut self, n: usize) {
78 self.max_idle_per_host = n;
79 }
80
81 pub async fn get(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
86 self.request("GET", url, headers, None).await
87 }
88
89 pub async fn head(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
94 self.request("HEAD", url, headers, None).await
95 }
96
97 pub async fn request(
111 &self,
112 method: &str,
113 url: &str,
114 headers: &[(&str, &str)],
115 body: Option<&[u8]>,
116 ) -> io::Result<Response> {
117 let (origin, path) = parse_url(url)?;
118
119 if let Some(conn) = self.take_pooled(&origin) {
121 match self
122 .try_once(conn, method, &origin, &path, headers, body)
123 .await
124 {
125 Ok(resp) => return Ok(resp),
126 Err(pooled_err) if !is_idempotent(method) => return Err(pooled_err),
130 Err(pooled_err) => {
131 let conn = Conn::connect(&origin, &self.tls).await?;
132 return self
133 .try_once(conn, method, &origin, &path, headers, body)
134 .await
135 .map_err(|retry_err| {
136 io::Error::new(
137 retry_err.kind(),
138 format!(
139 "{retry_err} (retry after stale pooled connection failed \
140 with: {pooled_err})"
141 ),
142 )
143 });
144 }
145 }
146 }
147
148 let conn = Conn::connect(&origin, &self.tls).await?;
149 self.try_once(conn, method, &origin, &path, headers, body)
150 .await
151 }
152
153 async fn try_once(
154 &self,
155 mut conn: Conn,
156 method: &str,
157 origin: &Origin,
158 path: &str,
159 headers: &[(&str, &str)],
160 body: Option<&[u8]>,
161 ) -> io::Result<Response> {
162 let is_head = method.eq_ignore_ascii_case("HEAD");
163 let host_header = origin.host_header();
164 let exchange = async {
165 write_request(&mut conn, method, &host_header, path, headers, body).await?;
166 read_response(&mut conn, is_head, self.max_response_bytes).await
167 };
168 let resp = match timeout(self.request_timeout, exchange).await {
169 Ok(result) => result?,
170 Err(_) => return Err(io::Error::new(io::ErrorKind::TimedOut, "request timed out")),
171 };
172 if resp.keep_alive {
173 self.put_pooled(origin, conn);
174 }
175 Ok(resp)
176 }
177
178 fn take_pooled(&self, origin: &Origin) -> Option<Conn> {
179 let mut pool = self.pool.lock().expect("http pool mutex poisoned");
180 pool.get_mut(origin).and_then(Vec::pop)
181 }
182
183 fn put_pooled(&self, origin: &Origin, conn: Conn) {
184 let mut pool = self.pool.lock().expect("http pool mutex poisoned");
185 let bucket = pool.entry(origin.clone()).or_default();
186 if bucket.len() < self.max_idle_per_host {
187 bucket.push(conn);
188 }
189 }
190}
191
192fn is_idempotent(method: &str) -> bool {
195 ["GET", "HEAD", "OPTIONS", "TRACE", "PUT", "DELETE"]
196 .iter()
197 .any(|m| method.eq_ignore_ascii_case(m))
198}
199
200fn parse_url(url: &str) -> io::Result<(Origin, String)> {
203 let uri = http::Uri::from_str(url)
204 .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, format!("bad URL: {e}")))?;
205 let secure = match uri.scheme_str() {
206 Some("https") => true,
207 Some("http") => false,
208 other => {
209 return Err(io::Error::new(
210 io::ErrorKind::InvalidInput,
211 format!("unsupported scheme: {other:?}"),
212 ))
213 }
214 };
215 let host = uri
216 .host()
217 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "URL missing host"))?
218 .to_string();
219 let port = uri.port_u16().unwrap_or(if secure { 443 } else { 80 });
220 let path = uri
221 .path_and_query()
222 .map(|pq| pq.as_str().to_string())
223 .unwrap_or_else(|| "/".to_string());
224 Ok((Origin { secure, host, port }, path))
225}
226
227#[cfg(test)]
228mod tests {
229 use super::*;
230 use std::io::{Read as _, Write as _};
231 use std::net::TcpListener;
232 use std::time::Instant;
233
234 fn read_request_head(stream: &mut std::net::TcpStream) {
236 let mut buf = Vec::new();
237 let mut byte = [0u8; 1];
238 while !buf.ends_with(b"\r\n\r\n") {
239 match stream.read(&mut byte) {
240 Ok(0) => break,
241 Ok(_) => buf.push(byte[0]),
242 Err(_) => break,
243 }
244 }
245 }
246
247 fn seed_stale_pooled_conn(client: &HttpClient, origin: &Origin, listener: &TcpListener) {
249 let conn = moirai::block_on(Conn::connect(origin, &client.tls))
250 .expect("pooled connection must establish");
251 let (stale, _) = listener.accept().expect("server must accept pooled conn");
252 drop(stale); client.put_pooled(origin, conn);
254 }
255
256 #[test]
257 fn idempotent_get_retries_stale_pooled_connection_and_succeeds() {
258 let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
259 let port = listener.local_addr().expect("addr").port();
260 let url = format!("http://127.0.0.1:{port}/x");
261
262 let client = HttpClient::new();
263 let origin = Origin {
264 secure: false,
265 host: "127.0.0.1".to_string(),
266 port,
267 };
268 seed_stale_pooled_conn(&client, &origin, &listener);
269
270 let server = std::thread::spawn(move || {
272 let (mut stream, _) = listener.accept().expect("retry connection must arrive");
273 read_request_head(&mut stream);
274 stream
275 .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok")
276 .expect("response write");
277 });
278
279 let resp = moirai::block_on(client.get(&url, &[]))
280 .expect("GET must transparently retry the stale pooled connection");
281 assert_eq!(resp.status, 200);
282 assert_eq!(resp.body, b"ok");
283 server.join().expect("server thread");
284 }
285
286 #[test]
287 fn non_idempotent_post_is_not_retried_after_stale_pooled_failure() {
288 let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
289 let port = listener.local_addr().expect("addr").port();
290 let url = format!("http://127.0.0.1:{port}/submit");
291
292 let client = HttpClient::new();
293 let origin = Origin {
294 secure: false,
295 host: "127.0.0.1".to_string(),
296 port,
297 };
298 seed_stale_pooled_conn(&client, &origin, &listener);
299
300 listener
302 .set_nonblocking(true)
303 .expect("nonblocking listener");
304
305 let err = moirai::block_on(client.request("POST", &url, &[], Some(b"payload")))
306 .expect_err("POST over a stale pooled connection must surface the error, not retry");
307 assert!(
312 matches!(
313 err.kind(),
314 io::ErrorKind::UnexpectedEof
315 | io::ErrorKind::ConnectionReset
316 | io::ErrorKind::ConnectionAborted
317 | io::ErrorKind::BrokenPipe
318 ),
319 "unexpected error kind {:?}: {err}",
320 err.kind()
321 );
322
323 let deadline = Instant::now() + Duration::from_millis(300);
325 while Instant::now() < deadline {
326 match listener.accept() {
327 Ok(_) => panic!("non-idempotent request must not be retried on a new connection"),
328 Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
329 std::thread::sleep(Duration::from_millis(10));
330 }
331 Err(e) => panic!("listener poll failed: {e}"),
332 }
333 }
334 }
335}