Skip to main content

reqwest_client/
reqwest_client.rs

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    /// Shared connection-management configuration for every client this type
30    /// builds. `read_timeout` sets an idle timeout on each body read (see
31    /// [`ReqwestClient::proxy_user_agent_and_read_timeout`]); `None` leaves
32    /// reads without a timeout.
33    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            // Detect and drop connections that have silently gone bad on a
38            // flaky path (NAT timeouts, resets) instead of reusing them. A
39            // stale reused HTTP/2 connection is a common source of
40            // `BadRecordMac` TLS errors against long-lived endpoints.
41            .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    /// Like [`ReqwestClient::proxy_and_user_agent`], but also applies a
71    /// per-read idle timeout. `read_timeout` fires only after that long with no
72    /// bytes received on a response body and resets on every chunk, so it
73    /// aborts a silently stalled stream without disturbing a healthy one that
74    /// merely goes quiet between chunks. Callers streaming long-lived responses
75    /// (e.g. LLM completions, which can pause for tens of seconds during
76    /// provider-side reasoning while keep-alive bytes still flow) should size
77    /// it comfortably above the provider's keep-alive interval.
78    ///
79    /// Note: on macOS the timeout is measured against a monotonic clock that
80    /// pauses during system sleep, so it does not fire from a suspend alone;
81    /// callers that need prompt detection of a connection killed while
82    /// suspended must re-validate the stream on wake themselves.
83    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            // Respect NO_PROXY env var
107            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            // Since we now have two executors, let's try to keep our footprint small
127            .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
149// This struct is essentially a re-implementation of
150// https://docs.rs/tokio-util/0.7.12/tokio_util/io/struct.ReaderStream.html
151// except outside of Tokio's aegis
152struct 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
211/// Implementation from <https://docs.rs/tokio-util/0.7.12/src/tokio_util/util/poll_buf.rs.html>
212/// Specialized for this use case
213fn 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        // Safety: `chunk_mut()` returns a `&mut UninitSlice`, and `UninitSlice` is a
226        // transparent wrapper around `[std::mem::MaybeUninit<u8>]`.
227        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        // SAFETY: Pin projection
231        let io_pin = unsafe { Pin::new_unchecked(io) };
232        // `futures::AsyncRead` reports the byte count as the poll's return
233        // value; `read_buf.filled()` stays empty because the reader writes
234        // through the initialized slice without advancing the `ReadBuf`.
235        std::task::ready!(io_pin.poll_read(cx, unfilled_portion)?)
236    };
237
238    // Safety: `initialize_unfilled()` zero-initialized the entire spare
239    // capacity, so the first `n` bytes are initialized no matter how many the
240    // reader actually wrote, and `advance_mut` panics rather than exceeding
241    // the capacity if `n` overstates the slice length.
242    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    /// Regression test: `StreamReader::poll_next` used to drop the reader it
337    /// `take()`s whenever the reader returned `Poll::Pending`, so the next
338    /// poll reported end-of-stream and streamed request bodies were silently
339    /// truncated. Readers backed by real I/O (e.g. `async_fs::File`) return
340    /// `Pending` on their very first read, so their uploads sent zero bytes.
341    #[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        // A reader that returns `Pending` before every chunk, like a reader
373        // backed by real I/O would.
374        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}