Skip to main content

rig_reqwest/
lib.rs

1#![cfg_attr(docsrs, feature(doc_cfg))]
2#![cfg_attr(
3    test,
4    allow(
5        clippy::expect_used,
6        clippy::indexing_slicing,
7        clippy::panic,
8        clippy::unwrap_used,
9        clippy::unreachable
10    )
11)]
12//! The bundled reqwest HTTP transport for Rig.
13//!
14//! [`ReqwestClient::default()`] is the process-wide client: built once on
15//! first use and shared by every clone, so every provider client made with
16//! it shares one connection pool. Construction never fails. When reqwest
17//! cannot build its client (a host with no CA store, say), every send
18//! reports that build failure in-band.
19//!
20//! Native requests and bodies enter the captured Tokio context on each poll,
21//! using a lazy fallback when no runtime is current. Callers retain ownership;
22//! dropping an operation cancels local work, not accepted remote work. Hosts
23//! must keep their runtime driven with I/O and timers enabled until operations
24//! finish. Missing drivers can panic; a stopped runtime causes I/O failure.
25//!
26//! ```no_run
27//! use rig_reqwest::ReqwestClient;
28//!
29//! let shared = ReqwestClient::default();
30//! let custom = ReqwestClient::from(rig_reqwest::reqwest::Client::new());
31//! # let _ = (shared, custom);
32//! ```
33
34pub use reqwest;
35
36/// A reqwest client implementing [`HttpClientExt`].
37///
38/// [`Default`] is the process-wide client, built once and shared by every
39/// clone; `From<reqwest::Client>` wraps a configured client and keeps its
40/// connection pool.
41#[derive(Clone, Debug)]
42pub struct ReqwestClient(Built);
43
44/// The process-wide client, erased behind [`DynHttpClient`].
45pub fn shared() -> DynHttpClient {
46    DynHttpClient::new(ReqwestClient::default())
47}
48
49/// What building the client produced: the client, or why reqwest refused.
50#[derive(Clone, Debug)]
51enum Built {
52    Client(Arc<reqwest::Client>),
53    /// The shared client could not be built: every send reports this, so
54    /// the failure surfaces where the first request is made rather than
55    /// where the client was named.
56    Failed(Arc<reqwest::Error>),
57}
58
59impl Default for ReqwestClient {
60    /// The process-wide client, built on first use. It never panics: a
61    /// client reqwest cannot build reports the build error on every send.
62    fn default() -> Self {
63        fn build() -> ReqwestClient {
64            ReqwestClient(match reqwest::Client::builder().build() {
65                Ok(client) => Built::Client(Arc::new(client)),
66                Err(error) => Built::Failed(Arc::new(error)),
67            })
68        }
69        #[cfg(not(target_family = "wasm"))]
70        {
71            static SHARED: std::sync::LazyLock<ReqwestClient> = std::sync::LazyLock::new(build);
72            SHARED.clone()
73        }
74        #[cfg(target_family = "wasm")]
75        {
76            thread_local! {
77                static SHARED: ReqwestClient = build();
78            }
79            SHARED.with(Clone::clone)
80        }
81    }
82}
83
84impl ReqwestClient {
85    /// The reqwest client, or `None` for the shared client when reqwest
86    /// could not build it.
87    pub fn inner(&self) -> Option<&reqwest::Client> {
88        match &self.0 {
89            Built::Client(client) => Some(client),
90            Built::Failed(_) => None,
91        }
92    }
93
94    /// Whether `self` and `other` are clones of one client.
95    #[cfg(test)]
96    fn same(&self, other: &Self) -> bool {
97        match (&self.0, &other.0) {
98            (Built::Client(a), Built::Client(b)) => Arc::ptr_eq(a, b),
99            (Built::Failed(a), Built::Failed(b)) => Arc::ptr_eq(a, b),
100            _ => false,
101        }
102    }
103}
104
105/// The error every send on an unbuilt client reports.
106fn unbuilt(error: &Arc<reqwest::Error>) -> Error {
107    Error::instance(TransportBuildError(Arc::clone(error)))
108}
109
110impl From<reqwest::Client> for ReqwestClient {
111    fn from(client: reqwest::Client) -> Self {
112        Self(Built::Client(Arc::new(client)))
113    }
114}
115
116/// A configured middleware client implementing [`HttpClientExt`].
117///
118/// Build the inner client with `reqwest_middleware::ClientBuilder`, then wrap
119/// it with `From<ClientWithMiddleware>`.
120#[cfg(any(
121    feature = "reqwest-middleware-rustls",
122    feature = "reqwest-middleware-native-tls"
123))]
124#[cfg_attr(
125    docsrs,
126    doc(cfg(any(
127        feature = "reqwest-middleware-rustls",
128        feature = "reqwest-middleware-native-tls"
129    )))
130)]
131#[derive(Clone, Debug)]
132pub struct ReqwestMiddlewareClient(reqwest_middleware::ClientWithMiddleware);
133
134#[cfg(any(
135    feature = "reqwest-middleware-rustls",
136    feature = "reqwest-middleware-native-tls"
137))]
138impl ReqwestMiddlewareClient {
139    /// Take the inner client back.
140    #[must_use]
141    pub fn into_inner(self) -> reqwest_middleware::ClientWithMiddleware {
142        self.0
143    }
144}
145
146#[cfg(any(
147    feature = "reqwest-middleware-rustls",
148    feature = "reqwest-middleware-native-tls"
149))]
150impl From<reqwest_middleware::ClientWithMiddleware> for ReqwestMiddlewareClient {
151    fn from(client: reqwest_middleware::ClientWithMiddleware) -> Self {
152        Self(client)
153    }
154}
155
156#[cfg(any(
157    feature = "reqwest-middleware-rustls",
158    feature = "reqwest-middleware-native-tls"
159))]
160impl AsRef<reqwest_middleware::ClientWithMiddleware> for ReqwestMiddlewareClient {
161    fn as_ref(&self) -> &reqwest_middleware::ClientWithMiddleware {
162        &self.0
163    }
164}
165
166#[cfg(not(target_family = "wasm"))]
167mod runtime;
168
169use bytes::Bytes;
170use futures::future::Either;
171use rig_http::http_client::{
172    DynHttpClient, Error, HttpClientExt, LazyBody, MultipartForm, Request, Response, Result,
173    StreamingResponse, multipart::PartContent,
174};
175use rig_http::wasm_compat::*;
176use std::pin::Pin;
177use std::sync::Arc;
178
179/// A transport build failure that displays the source chain and retains the
180/// original reqwest error as its source.
181#[derive(Debug)]
182struct TransportBuildError(Arc<reqwest::Error>);
183
184impl std::fmt::Display for TransportBuildError {
185    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
186        write!(
187            f,
188            "could not build the bundled reqwest transport: {}",
189            self.0
190        )?;
191        let mut source = std::error::Error::source(&*self.0);
192        while let Some(cause) = source {
193            write!(f, ": {cause}")?;
194            source = cause.source();
195        }
196        Ok(())
197    }
198}
199
200impl std::error::Error for TransportBuildError {
201    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
202        Some(&*self.0)
203    }
204}
205
206/// Wrap a reqwest transport error as [`Error::Instance`], retaining its source.
207///
208/// HTTP status failures are handled separately to preserve response headers
209/// and bodies.
210pub fn from_reqwest(err: reqwest::Error) -> Error {
211    Error::instance(err)
212}
213
214/// Read the status, headers and body off a failed `reqwest::Response` and
215/// build a non-success error that preserves the headers.
216async fn non_success_status_error(response: reqwest::Response) -> Error {
217    let status = response.status();
218    let headers = response.headers().clone();
219    let body = response
220        .text()
221        .await
222        .unwrap_or_else(|error| format!("failed to read error response body: {error}"));
223    Error::non_success_with_details(status, headers, body)
224}
225
226/// Keep successful bodies lazy and owned by the response on every executor.
227async fn into_response<U>(response: reqwest::Response) -> Result<Response<LazyBody<U>>>
228where
229    U: From<Bytes>,
230    U: WasmCompatSend + 'static,
231{
232    if !response.status().is_success() {
233        return Err(non_success_status_error(response).await);
234    }
235
236    let mut res = Response::builder().status(response.status());
237    if let Some(headers) = res.headers_mut() {
238        *headers = response.headers().clone();
239    }
240
241    let body = async {
242        let bytes = response.bytes().await.map_err(Error::instance)?;
243        Ok(U::from(bytes))
244    };
245    #[cfg(not(target_family = "wasm"))]
246    let body = runtime::bind(body)?;
247    let body: LazyBody<U> = Box::pin(body);
248    res.body(body).map_err(Error::Protocol)
249}
250
251fn streaming_head(response: &reqwest::Response) -> http::response::Builder {
252    #[cfg(not(target_family = "wasm"))]
253    let mut res = Response::builder()
254        .status(response.status())
255        .version(response.version());
256
257    #[cfg(target_family = "wasm")]
258    let mut res = Response::builder().status(response.status());
259
260    if let Some(hs) = res.headers_mut() {
261        *hs = response.headers().clone();
262    }
263    res
264}
265
266/// Convert an already-sent streaming response into the transport-agnostic
267/// [`StreamingResponse`], rejecting non-success statuses with the
268/// headers-preserving error. The body retains the request's reactor context.
269async fn into_streaming_response(response: reqwest::Response) -> Result<StreamingResponse> {
270    if !response.status().is_success() {
271        return Err(non_success_status_error(response).await);
272    }
273    let res = streaming_head(&response);
274
275    use futures::StreamExt;
276    let stream = response
277        .bytes_stream()
278        .map(|chunk| chunk.map_err(Error::instance));
279    #[cfg(not(target_family = "wasm"))]
280    let stream = runtime::bind_stream(stream)?;
281    let stream: Pin<Box<dyn WasmCompatSendStream<InnerItem = Result<Bytes>>>> = Box::pin(stream);
282    res.body(stream).map_err(Error::Protocol)
283}
284
285/// A part's content type was not a MIME type reqwest would accept.
286#[derive(Debug, thiserror::Error)]
287#[error("multipart part {part:?} has an unusable content type {content_type:?}: {source}")]
288struct InvalidPartContentType {
289    part: String,
290    content_type: String,
291    source: reqwest::Error,
292}
293
294/// Render a [`MultipartForm`] as a `reqwest::multipart::Form`.
295///
296/// Returns an error when a binary part has a content type reqwest rejects.
297pub fn multipart_form(value: MultipartForm) -> Result<reqwest::multipart::Form> {
298    let mut form = reqwest::multipart::Form::new();
299
300    for part in value.into_parts() {
301        let (name, content, filename, content_type) = part.into_pieces();
302        match content {
303            PartContent::Text(text) => {
304                form = form.text(name, text);
305            }
306            PartContent::Binary(bytes) => {
307                let mut req_part = reqwest::multipart::Part::bytes(bytes.to_vec());
308                if let Some(content_type) = content_type.as_ref() {
309                    req_part = req_part.mime_str(content_type.as_ref()).map_err(|source| {
310                        Error::instance(InvalidPartContentType {
311                            part: name.clone(),
312                            content_type: content_type.as_ref().to_string(),
313                            source,
314                        })
315                    })?;
316                }
317
318                if let Some(filename) = filename {
319                    req_part = req_part.file_name(filename);
320                }
321
322                form = form.part(name, req_part);
323            }
324        }
325    }
326
327    Ok(form)
328}
329
330/// Creates request builders for plain and middleware reqwest clients.
331trait ReqwestLike: Clone + WasmCompatSend + WasmCompatSync + 'static {
332    type Builder: RequestBuilderLike;
333    fn request_builder(&self, method: http::Method, url: String) -> Self::Builder;
334}
335
336trait RequestBuilderLike: Sized + WasmCompatSend + 'static {
337    fn with_headers(self, headers: http::HeaderMap) -> Self;
338    fn with_body(self, body: reqwest::Body) -> Self;
339    fn with_multipart(self, form: reqwest::multipart::Form) -> Self;
340    fn send_request(self) -> impl Future<Output = Result<reqwest::Response>> + WasmCompatSend;
341}
342
343impl ReqwestLike for Arc<reqwest::Client> {
344    type Builder = reqwest::RequestBuilder;
345    fn request_builder(&self, method: http::Method, url: String) -> Self::Builder {
346        self.request(method, url)
347    }
348}
349
350impl RequestBuilderLike for reqwest::RequestBuilder {
351    fn with_headers(self, headers: http::HeaderMap) -> Self {
352        self.headers(headers)
353    }
354    fn with_body(self, body: reqwest::Body) -> Self {
355        self.body(body)
356    }
357    fn with_multipart(self, form: reqwest::multipart::Form) -> Self {
358        self.multipart(form)
359    }
360    async fn send_request(self) -> Result<reqwest::Response> {
361        self.send().await.map_err(Error::instance)
362    }
363}
364
365#[cfg(any(
366    feature = "reqwest-middleware-rustls",
367    feature = "reqwest-middleware-native-tls"
368))]
369impl ReqwestLike for ReqwestMiddlewareClient {
370    type Builder = reqwest_middleware::RequestBuilder;
371    fn request_builder(&self, method: http::Method, url: String) -> Self::Builder {
372        self.0.request(method, url)
373    }
374}
375
376#[cfg(any(
377    feature = "reqwest-middleware-rustls",
378    feature = "reqwest-middleware-native-tls"
379))]
380impl RequestBuilderLike for reqwest_middleware::RequestBuilder {
381    fn with_headers(self, headers: http::HeaderMap) -> Self {
382        self.headers(headers)
383    }
384    fn with_body(self, body: reqwest::Body) -> Self {
385        self.body(body)
386    }
387    fn with_multipart(self, form: reqwest::multipart::Form) -> Self {
388        self.multipart(form)
389    }
390    async fn send_request(self) -> Result<reqwest::Response> {
391        self.send().await.map_err(Error::instance)
392    }
393}
394
395/// Select a reactor on first poll and retain it through response conversion.
396async fn drive<B, T, Convert, F>(request: B, convert: Convert) -> Result<T>
397where
398    B: RequestBuilderLike,
399    Convert: FnOnce(reqwest::Response) -> F + WasmCompatSend,
400    F: Future<Output = Result<T>> + WasmCompatSend,
401{
402    let operation = async move { convert(request.send_request().await?).await };
403    #[cfg(not(target_family = "wasm"))]
404    let operation = runtime::bind(operation)?;
405    operation.await
406}
407
408fn send_via<C, T, U>(
409    client: &C,
410    req: Request<T>,
411) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
412where
413    C: ReqwestLike,
414    T: Into<Bytes>,
415    U: From<Bytes> + WasmCompatSend + 'static,
416{
417    let (parts, body) = req.into_parts();
418    let req = client
419        .request_builder(parts.method, parts.uri.to_string())
420        .with_headers(parts.headers)
421        .with_body(body.into().into());
422
423    drive(req, into_response::<U>)
424}
425
426fn send_multipart_via<C, U>(
427    client: &C,
428    req: Request<MultipartForm>,
429) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
430where
431    C: ReqwestLike,
432    U: From<Bytes> + WasmCompatSend + 'static,
433{
434    let (parts, body) = req.into_parts();
435    // Reject invalid MIME types locally rather than sending incomplete metadata.
436    let form = multipart_form(body);
437    let req = form.map(|form| {
438        client
439            .request_builder(parts.method, parts.uri.to_string())
440            .with_headers(parts.headers)
441            .with_multipart(form)
442    });
443
444    async move { drive(req?, into_response::<U>).await }
445}
446
447fn send_streaming_via<C, T>(
448    client: &C,
449    req: Request<T>,
450) -> impl Future<Output = Result<StreamingResponse>> + WasmCompatSend
451where
452    C: ReqwestLike,
453    T: Into<Bytes> + WasmCompatSend,
454{
455    let (parts, body) = req.into_parts();
456    let req = client
457        .request_builder(parts.method, parts.uri.to_string())
458        .with_headers(parts.headers)
459        .with_body(body.into().into());
460
461    drive(req, into_streaming_response)
462}
463
464macro_rules! impl_http_client_ext_via {
465    ($(#[$attribute:meta])* $client:ty) => {
466        $(#[$attribute])*
467        impl HttpClientExt for $client {
468            fn send<T, U>(
469                &self,
470                req: Request<T>,
471            ) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
472            where
473                T: Into<Bytes>,
474                U: From<Bytes> + WasmCompatSend + 'static,
475            {
476                send_via(self, req)
477            }
478
479            fn send_multipart<U>(
480                &self,
481                req: Request<MultipartForm>,
482            ) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
483            where
484                U: From<Bytes> + WasmCompatSend + 'static,
485            {
486                send_multipart_via(self, req)
487            }
488
489            fn send_streaming<T>(
490                &self,
491                req: Request<T>,
492            ) -> impl Future<Output = Result<StreamingResponse>> + WasmCompatSend
493            where
494                T: Into<Bytes> + WasmCompatSend,
495            {
496                send_streaming_via(self, req)
497            }
498        }
499    };
500}
501
502impl HttpClientExt for ReqwestClient {
503    fn send<T, U>(
504        &self,
505        req: Request<T>,
506    ) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
507    where
508        T: Into<Bytes>,
509        U: From<Bytes> + WasmCompatSend + 'static,
510    {
511        match &self.0 {
512            Built::Client(client) => Either::Left(send_via(client, req)),
513            Built::Failed(error) => Either::Right(std::future::ready(Err(unbuilt(error)))),
514        }
515    }
516
517    fn send_multipart<U>(
518        &self,
519        req: Request<MultipartForm>,
520    ) -> impl Future<Output = Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
521    where
522        U: From<Bytes> + WasmCompatSend + 'static,
523    {
524        match &self.0 {
525            Built::Client(client) => Either::Left(send_multipart_via(client, req)),
526            Built::Failed(error) => Either::Right(std::future::ready(Err(unbuilt(error)))),
527        }
528    }
529
530    fn send_streaming<T>(
531        &self,
532        req: Request<T>,
533    ) -> impl Future<Output = Result<StreamingResponse>> + WasmCompatSend
534    where
535        T: Into<Bytes> + WasmCompatSend,
536    {
537        match &self.0 {
538            Built::Client(client) => Either::Left(send_streaming_via(client, req)),
539            Built::Failed(error) => Either::Right(std::future::ready(Err(unbuilt(error)))),
540        }
541    }
542}
543
544impl_http_client_ext_via!(
545    #[cfg(any(
546        feature = "reqwest-middleware-rustls",
547        feature = "reqwest-middleware-native-tls"
548    ))]
549    #[cfg_attr(
550        docsrs,
551        doc(cfg(any(
552            feature = "reqwest-middleware-rustls",
553            feature = "reqwest-middleware-native-tls"
554        )))
555    )]
556    ReqwestMiddlewareClient
557);
558
559// Compile-time thread-safety contract: the transport handle is shared across
560// threads by every host runtime.
561#[cfg(not(target_family = "wasm"))]
562const _: fn() = || {
563    fn assert_send_sync_static<T: Send + Sync + 'static>() {}
564    assert_send_sync_static::<ReqwestClient>();
565};
566
567#[cfg(test)]
568mod tests;