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)]
12pub use reqwest;
35
36#[derive(Clone, Debug)]
42pub struct ReqwestClient(Built);
43
44pub fn shared() -> DynHttpClient {
46 DynHttpClient::new(ReqwestClient::default())
47}
48
49#[derive(Clone, Debug)]
51enum Built {
52 Client(Arc<reqwest::Client>),
53 Failed(Arc<reqwest::Error>),
57}
58
59impl Default for ReqwestClient {
60 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 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 #[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
105fn 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#[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 #[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#[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
206pub fn from_reqwest(err: reqwest::Error) -> Error {
211 Error::instance(err)
212}
213
214async 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
226async 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
266async 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#[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
294pub 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
330trait 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
395async 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 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#[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;