1use std::error::Error;
2use std::sync::{LazyLock, OnceLock};
3use std::{borrow::Cow, mem, pin::Pin, task::Poll, time::Duration};
4
5use gpui_util::defer;
6
7use anyhow::anyhow;
8use bytes::{BufMut, Bytes, BytesMut};
9use futures::{AsyncRead, FutureExt as _, TryStreamExt as _};
10use http_client::{RedirectPolicy, RequestTimeout, Url, http};
11use regex::Regex;
12use reqwest::{
13 header::{HeaderMap, HeaderValue},
14 redirect,
15};
16
17const DEFAULT_CAPACITY: usize = 4096;
18static RUNTIME: OnceLock<tokio::runtime::Runtime> = OnceLock::new();
19static REDACT_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"key=[^&]+").unwrap());
20
21pub struct ReqwestClient {
22 client: reqwest::Client,
23 proxy: Option<Url>,
24 user_agent: Option<HeaderValue>,
25 handle: tokio::runtime::Handle,
26}
27
28impl ReqwestClient {
29 fn builder(read_timeout: Option<Duration>) -> reqwest::ClientBuilder {
34 let builder = reqwest::Client::builder()
35 .use_rustls_tls()
36 .connect_timeout(Duration::from_secs(10))
37 .tcp_keepalive(Duration::from_secs(30))
42 .pool_idle_timeout(Duration::from_secs(30))
43 .http2_keep_alive_interval(Duration::from_secs(15))
44 .http2_keep_alive_timeout(Duration::from_secs(10))
45 .http2_keep_alive_while_idle(true);
46 match read_timeout {
47 Some(read_timeout) => builder.read_timeout(read_timeout),
48 None => builder,
49 }
50 }
51
52 pub fn new() -> Self {
53 Self::builder(None)
54 .build()
55 .expect("Failed to initialize HTTP client")
56 .into()
57 }
58
59 pub fn user_agent(agent: &str) -> anyhow::Result<Self> {
60 let mut map = HeaderMap::new();
61 map.insert(http::header::USER_AGENT, HeaderValue::from_str(agent)?);
62 let client = Self::builder(None).default_headers(map).build()?;
63 Ok(client.into())
64 }
65
66 pub fn proxy_and_user_agent(proxy: Option<Url>, user_agent: &str) -> anyhow::Result<Self> {
67 Self::proxy_user_agent_and_read_timeout(proxy, user_agent, None)
68 }
69
70 pub fn proxy_user_agent_and_read_timeout(
84 proxy: Option<Url>,
85 user_agent: &str,
86 read_timeout: Option<Duration>,
87 ) -> anyhow::Result<Self> {
88 let user_agent = HeaderValue::from_str(user_agent)?;
89
90 let mut map = HeaderMap::new();
91 map.insert(http::header::USER_AGENT, user_agent.clone());
92 let mut client = Self::builder(read_timeout).default_headers(map);
93 let client_has_proxy;
94
95 if let Some(proxy) = proxy.as_ref().and_then(|proxy_url| {
96 reqwest::Proxy::all(proxy_url.clone())
97 .inspect_err(|e| {
98 log::error!(
99 "Failed to parse proxy URL '{}': {}",
100 proxy_url,
101 e.source().unwrap_or(&e as &_)
102 )
103 })
104 .ok()
105 }) {
106 client = client.proxy(proxy.no_proxy(reqwest::NoProxy::from_env()));
108 client_has_proxy = true;
109 } else {
110 client_has_proxy = false;
111 };
112
113 let client = client
114 .use_preconfigured_tls(http_client_tls::tls_config())
115 .build()?;
116 let mut client: ReqwestClient = client.into();
117 client.proxy = client_has_proxy.then_some(proxy).flatten();
118 client.user_agent = Some(user_agent);
119 Ok(client)
120 }
121}
122
123pub fn runtime() -> &'static tokio::runtime::Runtime {
124 RUNTIME.get_or_init(|| {
125 tokio::runtime::Builder::new_multi_thread()
126 .worker_threads(1)
128 .enable_all()
129 .build()
130 .expect("Failed to initialize HTTP client")
131 })
132}
133
134impl From<reqwest::Client> for ReqwestClient {
135 fn from(client: reqwest::Client) -> Self {
136 let handle = tokio::runtime::Handle::try_current().unwrap_or_else(|_| {
137 log::debug!("no tokio runtime found, creating one for Reqwest...");
138 runtime().handle().clone()
139 });
140 Self {
141 client,
142 handle,
143 proxy: None,
144 user_agent: None,
145 }
146 }
147}
148
149struct StreamReader {
153 reader: Option<Pin<Box<dyn futures::AsyncRead + Send + Sync>>>,
154 buf: BytesMut,
155 capacity: usize,
156}
157
158impl StreamReader {
159 fn new(reader: Pin<Box<dyn futures::AsyncRead + Send + Sync>>) -> Self {
160 Self {
161 reader: Some(reader),
162 buf: BytesMut::new(),
163 capacity: DEFAULT_CAPACITY,
164 }
165 }
166}
167
168impl futures::Stream for StreamReader {
169 type Item = std::io::Result<Bytes>;
170
171 fn poll_next(
172 mut self: Pin<&mut Self>,
173 cx: &mut std::task::Context<'_>,
174 ) -> Poll<Option<Self::Item>> {
175 let mut this = self.as_mut();
176
177 let mut reader = match this.reader.take() {
178 Some(r) => r,
179 None => return Poll::Ready(None),
180 };
181
182 if this.buf.capacity() == 0 {
183 let capacity = this.capacity;
184 this.buf.reserve(capacity);
185 }
186
187 match poll_read_buf(&mut reader, cx, &mut this.buf) {
188 Poll::Pending => {
189 self.reader = Some(reader);
190
191 Poll::Pending
192 }
193 Poll::Ready(Err(err)) => {
194 self.reader = None;
195
196 Poll::Ready(Some(Err(err)))
197 }
198 Poll::Ready(Ok(0)) => {
199 self.reader = None;
200 Poll::Ready(None)
201 }
202 Poll::Ready(Ok(_)) => {
203 let chunk = this.buf.split();
204 self.reader = Some(reader);
205 Poll::Ready(Some(Ok(chunk.freeze())))
206 }
207 }
208 }
209}
210
211fn poll_read_buf(
214 io: &mut Pin<Box<dyn futures::AsyncRead + Send + Sync>>,
215 cx: &mut std::task::Context<'_>,
216 buf: &mut BytesMut,
217) -> Poll<std::io::Result<usize>> {
218 if !buf.has_remaining_mut() {
219 return Poll::Ready(Ok(0));
220 }
221
222 let n = {
223 let dst = buf.chunk_mut();
224
225 let dst = unsafe { &mut *(dst as *mut _ as *mut [std::mem::MaybeUninit<u8>]) };
228 let mut read_buf = tokio::io::ReadBuf::uninit(dst);
229 let unfilled_portion = read_buf.initialize_unfilled();
230 let io_pin = unsafe { Pin::new_unchecked(io) };
232 std::task::ready!(io_pin.poll_read(cx, unfilled_portion)?)
236 };
237
238 unsafe {
243 buf.advance_mut(n);
244 }
245
246 Poll::Ready(Ok(n))
247}
248
249fn redact_error(mut error: reqwest::Error) -> reqwest::Error {
250 if let Some(url) = error.url_mut()
251 && let Some(query) = url.query()
252 && let Cow::Owned(redacted) = REDACT_REGEX.replace_all(query, "key=REDACTED")
253 {
254 url.set_query(Some(redacted.as_str()));
255 }
256 error
257}
258
259impl http_client::HttpClient for ReqwestClient {
260 fn proxy(&self) -> Option<&Url> {
261 self.proxy.as_ref()
262 }
263
264 fn user_agent(&self) -> Option<&HeaderValue> {
265 self.user_agent.as_ref()
266 }
267
268 fn send(
269 &self,
270 req: http::Request<http_client::AsyncBody>,
271 ) -> futures::future::BoxFuture<
272 'static,
273 anyhow::Result<http_client::Response<http_client::AsyncBody>>,
274 > {
275 let (parts, body) = req.into_parts();
276
277 let mut request_builder = self.client.request(parts.method, parts.uri.to_string());
278 request_builder = request_builder.headers(parts.headers);
279 if let Some(redirect_policy) = parts.extensions.get::<RedirectPolicy>() {
280 request_builder = request_builder.redirect_policy(match redirect_policy {
281 RedirectPolicy::NoFollow => redirect::Policy::none(),
282 RedirectPolicy::FollowLimit(limit) => redirect::Policy::limited(*limit as usize),
283 RedirectPolicy::FollowAll => redirect::Policy::limited(100),
284 });
285 }
286 if let Some(timeout) = parts.extensions.get::<RequestTimeout>() {
287 request_builder = request_builder.timeout(timeout.0);
288 }
289 let request = request_builder.body(match body.0 {
290 http_client::Inner::Empty => reqwest::Body::default(),
291 http_client::Inner::Bytes(cursor) => cursor.into_inner().into(),
292 http_client::Inner::AsyncReader(stream) => {
293 reqwest::Body::wrap_stream(StreamReader::new(stream))
294 }
295 });
296
297 let handle = self.handle.clone();
298 async move {
299 let join_handle = handle.spawn(async { request.send().await });
300 let abort_handle = join_handle.abort_handle();
301 let _abort_on_drop = defer(move || abort_handle.abort());
302
303 let mut response = join_handle.await?.map_err(redact_error)?;
304
305 let headers = mem::take(response.headers_mut());
306 let mut builder = http::Response::builder()
307 .status(response.status().as_u16())
308 .version(response.version());
309 *builder.headers_mut().unwrap() = headers;
310
311 let bytes = response
312 .bytes_stream()
313 .map_err(futures::io::Error::other)
314 .into_async_read();
315 let body = http_client::AsyncBody::from_reader(bytes);
316
317 builder.body(body).map_err(|e| anyhow!(e))
318 }
319 .boxed()
320 }
321}
322
323#[cfg(test)]
324mod tests {
325 use std::io::{BufRead as _, BufReader, Read as _, Write as _};
326 use std::net::TcpListener;
327 use std::time::{Duration, Instant};
328
329 use futures::AsyncReadExt as _;
330 use http_client::{
331 AsyncBody, HttpClient, HttpRequestExt as _, Method, Request as HttpRequest, Url,
332 };
333
334 use crate::ReqwestClient;
335
336 #[test]
342 fn test_streamed_body_survives_pending_reader() {
343 let payload: Vec<u8> = (0..30_000usize).map(|byte| (byte % 251) as u8).collect();
344
345 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
346 let address = listener.local_addr().unwrap();
347 let expected_payload = payload.clone();
348 let server = std::thread::spawn(move || {
349 let (mut stream, _) = listener.accept().unwrap();
350 let mut request = Vec::new();
351 let mut buffer = [0u8; 8192];
352 loop {
353 let read = stream.read(&mut buffer).unwrap();
354 assert_ne!(read, 0, "client closed the connection mid-request");
355 request.extend_from_slice(&buffer[..read]);
356 if let Some(position) = request.windows(4).position(|w| w == b"\r\n\r\n") {
357 let body_start = position + 4;
358 while request.len() - body_start < expected_payload.len() {
359 let read = stream.read(&mut buffer).unwrap();
360 assert_ne!(read, 0, "client closed the connection mid-body");
361 request.extend_from_slice(&buffer[..read]);
362 }
363 assert_eq!(&request[body_start..], &expected_payload);
364 break;
365 }
366 }
367 stream
368 .write_all(b"HTTP/1.1 200 OK\r\ncontent-length: 0\r\nconnection: close\r\n\r\n")
369 .unwrap();
370 });
371
372 struct PendingFirstReader {
375 data: std::io::Cursor<Vec<u8>>,
376 ready: bool,
377 }
378
379 impl futures::AsyncRead for PendingFirstReader {
380 fn poll_read(
381 mut self: std::pin::Pin<&mut Self>,
382 cx: &mut std::task::Context<'_>,
383 buf: &mut [u8],
384 ) -> std::task::Poll<std::io::Result<usize>> {
385 if self.ready {
386 self.ready = false;
387 std::task::Poll::Ready(self.data.read(buf))
388 } else {
389 self.ready = true;
390 cx.waker().wake_by_ref();
391 std::task::Poll::Pending
392 }
393 }
394 }
395
396 let reader = PendingFirstReader {
397 data: std::io::Cursor::new(payload.clone()),
398 ready: false,
399 };
400
401 let client = ReqwestClient::new();
402 let request = HttpRequest::builder()
403 .method(Method::PUT)
404 .uri(format!("http://{address}/upload"))
405 .header("Content-Length", payload.len().to_string())
406 .body(AsyncBody::from_reader(reader))
407 .unwrap();
408 let response = futures::executor::block_on(client.send(request)).unwrap();
409 assert!(response.status().is_success());
410 server.join().unwrap();
411 }
412
413 #[test]
414 fn test_request_timeout_applies_while_reading_response_body() {
415 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
416 let address = listener.local_addr().unwrap();
417 let server = std::thread::spawn(move || {
418 let (mut stream, _) = listener.accept().unwrap();
419 let mut reader = BufReader::new(&mut stream);
420 let mut line = String::new();
421 loop {
422 line.clear();
423 assert_ne!(reader.read_line(&mut line).unwrap(), 0);
424 if line == "\r\n" {
425 break;
426 }
427 }
428 drop(reader);
429 stream
430 .write_all(b"HTTP/1.1 200 OK\r\ncontent-length: 1\r\n\r\n")
431 .unwrap();
432 stream
433 .set_read_timeout(Some(Duration::from_secs(1)))
434 .unwrap();
435 let mut buffer = [0; 1];
436 match stream.read(&mut buffer) {
437 Ok(_) => {}
438 Err(error)
439 if matches!(
440 error.kind(),
441 std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock
442 ) => {}
443 Err(error) => panic!("failed while waiting for the client to close: {error}"),
444 }
445 });
446
447 let client = ReqwestClient::new();
448 let request = HttpRequest::get(format!("http://{address}/"))
449 .timeout(Duration::from_millis(100))
450 .body(AsyncBody::default())
451 .unwrap();
452 let started_at = Instant::now();
453 let mut response = futures::executor::block_on(client.send(request)).unwrap();
454 let mut body = Vec::new();
455 let result = futures::executor::block_on(response.body_mut().read_to_end(&mut body));
456 assert!(result.is_err(), "the response body should time out");
457 assert!(
458 started_at.elapsed() < Duration::from_millis(500),
459 "the request timeout should not wait for the server to close the connection"
460 );
461 drop(response);
462 server.join().unwrap();
463 }
464
465 #[test]
466 fn test_proxy_uri() {
467 let client = ReqwestClient::new();
468 assert_eq!(client.proxy(), None);
469
470 let proxy = Url::parse("http://localhost:10809").unwrap();
471 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
472 assert_eq!(client.proxy(), Some(&proxy));
473
474 let proxy = Url::parse("https://localhost:10809").unwrap();
475 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
476 assert_eq!(client.proxy(), Some(&proxy));
477
478 let proxy = Url::parse("socks4://localhost:10808").unwrap();
479 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
480 assert_eq!(client.proxy(), Some(&proxy));
481
482 let proxy = Url::parse("socks4a://localhost:10808").unwrap();
483 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
484 assert_eq!(client.proxy(), Some(&proxy));
485
486 let proxy = Url::parse("socks5://localhost:10808").unwrap();
487 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
488 assert_eq!(client.proxy(), Some(&proxy));
489
490 let proxy = Url::parse("socks5h://localhost:10808").unwrap();
491 let client = ReqwestClient::proxy_and_user_agent(Some(proxy.clone()), "test").unwrap();
492 assert_eq!(client.proxy(), Some(&proxy));
493 }
494
495 #[test]
496 fn test_invalid_proxy_uri() {
497 let proxy = Url::parse("socks://127.0.0.1:20170").unwrap();
498 let client = ReqwestClient::proxy_and_user_agent(Some(proxy), "test").unwrap();
499 assert!(
500 client.proxy.is_none(),
501 "An invalid proxy URL should add no proxy to the client!"
502 )
503 }
504}