1use blitz_traits::net::{AbortSignal, Body, Bytes, NetHandler, NetProvider, NetWaker, Request};
17use blitz_traits::platform::{
18 FetchError, FetchHandler, FetchProvider, FetchRequest, FetchResponse, HeaderMap, StatusCode,
19};
20use data_url::DataUrl;
21use std::{
22 collections::HashMap,
23 marker::PhantomData,
24 pin::Pin,
25 sync::{Arc, Mutex},
26 task::Poll,
27};
28use tokio::sync::Semaphore;
29
30#[cfg(feature = "cache")]
31use http_cache_reqwest::{
32 CACacheManager, Cache, CacheMode, CacheOptions, HttpCache, HttpCacheOptions,
33};
34
35const USER_AGENT: &str = "Mozilla/5.0 (X11; Linux x86_64; rv:60.0) Gecko/20100101 Firefox/81.0";
36
37const PER_HOST_MAX_CONCURRENT: usize = 6;
39
40type HostLimits = Arc<Mutex<HashMap<String, Arc<Semaphore>>>>;
41
42#[cfg(feature = "cache")]
43type Client = reqwest_middleware::ClientWithMiddleware;
44#[cfg(not(feature = "cache"))]
45type Client = reqwest::Client;
46
47#[cfg(feature = "cache")]
48type RequestBuilder = reqwest_middleware::RequestBuilder;
49#[cfg(not(feature = "cache"))]
50type RequestBuilder = reqwest::RequestBuilder;
51
52#[cfg(feature = "cache")]
53fn get_cache_path() -> std::path::PathBuf {
54 use directories::ProjectDirs;
55 let path = ProjectDirs::from("com", "DioxusLabs", "Blitz")
56 .expect("Failed to find cache directory")
57 .cache_dir()
58 .to_owned();
59 #[cfg(feature = "tracing")]
60 tracing::info!(path = ?path.display(), "Using cache dir");
61 path
62}
63
64#[cfg(target_arch = "wasm32")]
65fn spawn(fut: impl Future + 'static) {
66 wasm_bindgen_futures::spawn_local(async move {
67 fut.await;
68 });
69}
70
71#[cfg(not(target_arch = "wasm32"))]
72fn spawn<F>(fut: F)
73where
74 F: Future + Send + 'static,
75 F::Output: Send + 'static,
76{
77 tokio::spawn(fut);
78}
79
80pub struct Provider {
81 client: Client,
82 waker: Arc<dyn NetWaker>,
83 per_host_limits: HostLimits,
84 #[cfg(feature = "cache")]
85 cache_manager: CACacheManager,
86}
87impl Provider {
88 pub fn new(waker: Option<Arc<dyn NetWaker>>) -> Self {
89 let builder = reqwest::Client::builder();
90 #[cfg(feature = "cookies")]
91 let builder = builder.cookie_store(true);
92 let client = builder.build().unwrap();
93
94 #[cfg(feature = "cache")]
95 let cache_manager = CACacheManager::new(get_cache_path(), true);
96
97 #[cfg(feature = "cache")]
98 let client = reqwest_middleware::ClientBuilder::new(client)
99 .with(Cache(HttpCache {
100 mode: CacheMode::Default,
101 manager: cache_manager.clone(),
102 options: HttpCacheOptions {
103 cache_options: Some(CacheOptions {
112 shared: false,
113 ..Default::default()
114 }),
115 ..Default::default()
116 },
117 }))
118 .build();
119
120 let waker = waker.unwrap_or(Arc::new(DummyNetWaker));
121 Self {
122 client,
123 waker,
124 per_host_limits: Arc::new(Mutex::new(HashMap::new())),
125 #[cfg(feature = "cache")]
126 cache_manager,
127 }
128 }
129 pub fn shared(waker: Option<Arc<dyn NetWaker>>) -> Arc<dyn NetProvider> {
130 Arc::new(Self::new(waker))
131 }
132 pub fn is_empty(&self) -> bool {
133 Arc::strong_count(&self.waker) == 1
134 }
135 pub fn count(&self) -> usize {
136 Arc::strong_count(&self.waker) - 1
137 }
138
139 #[cfg(feature = "cache")]
140 pub async fn clear_cache(&self) {
141 if let Err(e) = self.cache_manager.clear().await {
142 #[cfg(feature = "tracing")]
143 tracing::error!("Failed to clear HTTP cache: {:?}", e);
144 #[cfg(not(feature = "tracing"))]
145 let _ = e;
146 }
147 }
148}
149impl Provider {
150 async fn fetch_inner(
151 client: Client,
152 request: Request,
153 per_host_limits: HostLimits,
154 ) -> Result<(String, Bytes), ProviderError> {
155 match request.url.scheme() {
156 "data" => {
157 let data_url = DataUrl::process(request.url.as_str())?;
158 let decoded = data_url.decode_to_vec()?;
159 Ok((request.url.to_string(), Bytes::from(decoded.0)))
160 }
161 "file" => {
162 let file_content = std::fs::read(request.url.path())?;
163 Ok((request.url.to_string(), Bytes::from(file_content)))
164 }
165 _ => Self::fetch_http(client, request, per_host_limits).await,
166 }
167 }
168
169 async fn fetch_http(
170 client: Client,
171 request: Request,
172 per_host_limits: HostLimits,
173 ) -> Result<(String, Bytes), ProviderError> {
174 let host_key = request
177 .url
178 .host_str()
179 .map(str::to_owned)
180 .unwrap_or_default();
181 let semaphore = {
182 let mut map = per_host_limits.lock().unwrap();
183 map.entry(host_key)
184 .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
185 .clone()
186 };
187 let _permit = semaphore
188 .acquire()
189 .await
190 .expect("per-host semaphore was closed");
191
192 let mut req = client
193 .request(request.method, request.url)
194 .headers(request.headers)
195 .header("User-Agent", USER_AGENT);
196
197 if let Some(content_type) = request.content_type.as_ref() {
198 req = req.header("Content-Type", content_type);
199 }
200
201 let req = req
202 .apply_body(request.body, request.content_type.as_deref())
203 .await;
204 let response = req.send().await?;
205 let status = response.status();
206 let final_url = response.url().to_string();
207
208 if status.is_success() {
209 return Ok((final_url, response.bytes().await?));
210 }
211
212 #[cfg(feature = "tracing")]
213 tracing::warn!(
214 url = final_url.as_str(),
215 status = status.as_u16(),
216 "HTTP error status"
217 );
218 Err(ProviderError::HttpStatus {
219 status,
220 url: final_url,
221 })
222 }
223
224 #[allow(clippy::type_complexity)]
225 pub fn fetch_with_callback(
226 &self,
227 request: Request,
228 callback: Box<dyn FnOnce(Result<(String, Bytes), ProviderError>) + Send + Sync + 'static>,
229 ) {
230 #[cfg(feature = "tracing")]
231 let url = request.url.to_string();
232
233 let client = self.client.clone();
234 let per_host_limits = self.per_host_limits.clone();
235 spawn(async move {
236 let result = Self::fetch_inner(client, request, per_host_limits).await;
237
238 #[cfg(feature = "tracing")]
239 if let Err(e) = &result {
240 #[cfg(feature = "tracing")]
241 tracing::error!(url = url.as_str(), error = ?e, "Fetching");
242 } else {
243 #[cfg(feature = "tracing")]
244 tracing::info!(url = url.as_str(), "Success fetching");
245 }
246
247 callback(result);
248 });
249 }
250
251 pub async fn fetch_async(&self, request: Request) -> Result<(String, Bytes), ProviderError> {
252 #[cfg(feature = "tracing")]
253 let url = request.url.to_string();
254
255 let client = self.client.clone();
256 let per_host_limits = self.per_host_limits.clone();
257 let result = Self::fetch_inner(client, request, per_host_limits).await;
258
259 #[cfg(feature = "tracing")]
260 if let Err(e) = &result {
261 #[cfg(feature = "tracing")]
262 tracing::error!(url = url.as_str(), error = ?e, "Fetching");
263 } else {
264 #[cfg(feature = "tracing")]
265 tracing::info!(url = url.as_str(), "Success fetching");
266 }
267
268 result
269 }
270
271 pub async fn fetch_response_async(
296 &self,
297 request: Request,
298 ) -> Result<FetchResponse, ProviderError> {
299 let url = request.url.clone();
300 match url.scheme() {
301 "data" => {
302 let (body, headers) = {
305 let data_url = DataUrl::process(url.as_str())?;
306 let decoded = data_url.decode_to_vec()?;
307 let mut headers = HeaderMap::new();
308 if let Ok(value) = data_url.mime_type().to_string().parse() {
309 headers.insert(blitz_traits::platform::http::header::CONTENT_TYPE, value);
310 }
311 (Bytes::from(decoded.0), headers)
312 };
313 Ok(FetchResponse::new(url, StatusCode::OK)
314 .headers(headers)
315 .body(body))
316 }
317 "file" => {
318 let file_content = std::fs::read(url.path())?;
319 Ok(FetchResponse::new(url, StatusCode::OK).body(Bytes::from(file_content)))
320 }
321 _ => {
322 let client = self.client.clone();
323 let per_host_limits = self.per_host_limits.clone();
324 Self::fetch_http_response(client, request, per_host_limits).await
325 }
326 }
327 }
328
329 async fn fetch_http_response(
337 client: Client,
338 request: Request,
339 per_host_limits: HostLimits,
340 ) -> Result<FetchResponse, ProviderError> {
341 let host_key = request
342 .url
343 .host_str()
344 .map(str::to_owned)
345 .unwrap_or_default();
346 let semaphore = {
347 let mut map = per_host_limits.lock().unwrap();
348 map.entry(host_key)
349 .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
350 .clone()
351 };
352 let _permit = semaphore
353 .acquire()
354 .await
355 .expect("per-host semaphore was closed");
356
357 let mut req = client
358 .request(request.method, request.url)
359 .headers(request.headers)
360 .header("User-Agent", USER_AGENT);
361
362 if let Some(content_type) = request.content_type.as_ref() {
363 req = req.header("Content-Type", content_type);
364 }
365
366 let req = req
367 .apply_body(request.body, request.content_type.as_deref())
368 .await;
369 let response = req.send().await?;
370 let status = response.status();
371 let final_url = response.url().clone();
372
373 if !status.is_success() {
374 #[cfg(feature = "tracing")]
375 tracing::warn!(
376 url = final_url.as_str(),
377 status = status.as_u16(),
378 "HTTP error status"
379 );
380 return Err(ProviderError::HttpStatus {
381 status,
382 url: final_url.to_string(),
383 });
384 }
385
386 let headers = response.headers().clone();
388 Ok(FetchResponse::new(final_url, status)
389 .headers(headers)
390 .body(response.bytes().await?))
391 }
392}
393
394impl Provider {
402 async fn platform_fetch_inner(
403 client: Client,
404 request: FetchRequest,
405 per_host_limits: HostLimits,
406 ) -> Result<FetchResponse, FetchError> {
407 match request.url.scheme() {
408 "data" => Self::platform_fetch_data(request),
409 "file" => Self::platform_fetch_file(request),
410 "http" | "https" => Self::platform_fetch_http(client, request, per_host_limits).await,
411 scheme => Err(FetchError::UnsupportedScheme(scheme.to_owned())),
412 }
413 }
414
415 fn platform_fetch_data(request: FetchRequest) -> Result<FetchResponse, FetchError> {
422 let data_url = DataUrl::process(request.url.as_str())
423 .map_err(|err| FetchError::InvalidRequest(format!("{err:?}")))?;
424 let mime = data_url.mime_type().to_string();
425 let (body, _) = data_url
426 .decode_to_vec()
427 .map_err(|err| FetchError::InvalidRequest(format!("{err:?}")))?;
428
429 let mut headers = HeaderMap::new();
430 if let Ok(value) = mime.parse() {
431 headers.insert(reqwest::header::CONTENT_TYPE, value);
434 }
435
436 Ok(FetchResponse::new(request.url, StatusCode::OK)
437 .headers(headers)
438 .body(Bytes::from(body)))
439 }
440
441 fn platform_fetch_file(request: FetchRequest) -> Result<FetchResponse, FetchError> {
455 let path = request.url.to_file_path().map_err(|()| {
456 FetchError::InvalidRequest(format!("not a local path: {}", request.url))
457 })?;
458
459 let body = std::fs::read(path).map_err(|err| FetchError::Network(err.to_string()))?;
460
461 Ok(FetchResponse::new(request.url, StatusCode::OK).body(Bytes::from(body)))
462 }
463
464 async fn platform_fetch_http(
465 client: Client,
466 request: FetchRequest,
467 per_host_limits: HostLimits,
468 ) -> Result<FetchResponse, FetchError> {
469 let host_key = request
472 .url
473 .host_str()
474 .map(str::to_owned)
475 .unwrap_or_default();
476 let semaphore = {
477 let mut map = per_host_limits.lock().unwrap();
478 map.entry(host_key)
479 .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
480 .clone()
481 };
482 let _permit = semaphore
483 .acquire()
484 .await
485 .expect("per-host semaphore was closed");
486
487 let mut req = client
488 .request(request.method, request.url)
489 .headers(request.headers)
490 .header("User-Agent", USER_AGENT);
491
492 if let Some(body) = request.body {
493 req = req.body(body);
494 }
495
496 let response = req
497 .send()
498 .await
499 .map_err(|err| FetchError::Network(err.to_string()))?;
500
501 let status = response.status();
503 let headers = response.headers().clone();
504 let url = response.url().clone();
505 let body = response
506 .bytes()
507 .await
508 .map_err(|err| FetchError::Network(err.to_string()))?;
509
510 Ok(FetchResponse::new(url, status).headers(headers).body(body))
511 }
512}
513
514impl FetchProvider for Provider {
515 fn fetch(&self, request: FetchRequest, handler: Box<dyn FetchHandler>) {
516 let client = self.client.clone();
517 let per_host_limits = self.per_host_limits.clone();
518
519 #[cfg(feature = "tracing")]
520 let url = request.url.to_string();
521
522 spawn(async move {
523 let result = Self::platform_fetch_inner(client, request, per_host_limits).await;
524
525 #[cfg(feature = "tracing")]
526 match &result {
527 Ok(response) => tracing::info!(
528 url = url.as_str(),
529 status = response.status.as_u16(),
530 "fetch complete"
531 ),
532 Err(error) => tracing::error!(url = url.as_str(), error = ?error, "fetch failed"),
533 }
534
535 handler.complete(result);
536 });
537 }
538}
539
540impl NetProvider for Provider {
541 fn fetch(&self, doc_id: usize, mut request: Request, handler: Box<dyn NetHandler>) {
542 let client = self.client.clone();
543 let per_host_limits = self.per_host_limits.clone();
544
545 #[cfg(feature = "tracing")]
546 tracing::info!(url = request.url.as_str(), "Fetching");
547
548 let waker = self.waker.clone();
549 spawn(async move {
550 #[cfg(feature = "tracing")]
551 let url = request.url.to_string();
552
553 let signal = request.signal.take();
554 let result = if let Some(signal) = signal {
555 AbortFetch::new(
556 signal,
557 Box::pin(
558 async move { Self::fetch_inner(client, request, per_host_limits).await },
559 ),
560 )
561 .await
562 } else {
563 Self::fetch_inner(client, request, per_host_limits).await
564 };
565
566 waker.wake(doc_id);
567
568 match result {
569 Ok((response_url, bytes)) => {
570 handler.bytes(response_url, bytes);
571 #[cfg(feature = "tracing")]
572 tracing::info!(url = url.as_str(), "Success fetching");
573 }
574 Err(e) => {
575 #[cfg(feature = "tracing")]
576 tracing::error!(url = url.as_str(), error = ?e, "Error fetching");
577 #[cfg(not(feature = "tracing"))]
578 let _ = e;
579 }
580 };
581 });
582 }
583}
584
585struct AbortFetch<F, T> {
586 signal: AbortSignal,
587 future: F,
588 _rt: PhantomData<T>,
589}
590
591impl<F, T> AbortFetch<F, T> {
592 fn new(signal: AbortSignal, future: F) -> Self {
593 Self {
594 signal,
595 future,
596 _rt: PhantomData,
597 }
598 }
599}
600
601impl<F, T> Future for AbortFetch<F, T>
602where
603 F: Future + Unpin + 'static,
604 F::Output: Into<Result<T, ProviderError>> + 'static,
605 T: Unpin,
606{
607 type Output = Result<T, ProviderError>;
608
609 fn poll(
610 mut self: std::pin::Pin<&mut Self>,
611 cx: &mut std::task::Context<'_>,
612 ) -> std::task::Poll<Self::Output> {
613 if self.signal.aborted() {
614 return Poll::Ready(Err(ProviderError::Abort));
615 }
616
617 match Pin::new(&mut self.future).poll(cx) {
618 Poll::Ready(output) => Poll::Ready(output.into()),
619 Poll::Pending => Poll::Pending,
620 }
621 }
622}
623
624#[derive(Debug)]
625pub enum ProviderError {
626 Abort,
627 Io(std::io::Error),
628 DataUrl(data_url::DataUrlError),
629 DataUrlBase64(data_url::forgiving_base64::InvalidBase64),
630 ReqwestError(reqwest::Error),
631 #[cfg(feature = "cache")]
632 ReqwestMiddlewareError(reqwest_middleware::Error),
633 HttpStatus {
634 status: reqwest::StatusCode,
635 url: String,
636 },
637}
638
639impl std::fmt::Display for ProviderError {
640 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
641 match self {
642 Self::Abort => write!(f, "request aborted"),
643 Self::Io(e) => write!(f, "io error: {e}"),
644 Self::DataUrl(e) => write!(f, "data url error: {e:?}"),
645 Self::DataUrlBase64(e) => write!(f, "data url base64 error: {e:?}"),
646 Self::ReqwestError(e) => write!(f, "reqwest error: {e}"),
647 #[cfg(feature = "cache")]
648 Self::ReqwestMiddlewareError(e) => write!(f, "reqwest middleware error: {e}"),
649 Self::HttpStatus { status, url } => write!(f, "HTTP {status} for {url}"),
650 }
651 }
652}
653
654impl From<std::io::Error> for ProviderError {
655 fn from(value: std::io::Error) -> Self {
656 Self::Io(value)
657 }
658}
659
660impl From<data_url::DataUrlError> for ProviderError {
661 fn from(value: data_url::DataUrlError) -> Self {
662 Self::DataUrl(value)
663 }
664}
665
666impl From<data_url::forgiving_base64::InvalidBase64> for ProviderError {
667 fn from(value: data_url::forgiving_base64::InvalidBase64) -> Self {
668 Self::DataUrlBase64(value)
669 }
670}
671
672impl From<reqwest::Error> for ProviderError {
673 fn from(value: reqwest::Error) -> Self {
674 Self::ReqwestError(value)
675 }
676}
677
678#[cfg(feature = "cache")]
679impl From<reqwest_middleware::Error> for ProviderError {
680 fn from(value: reqwest_middleware::Error) -> Self {
681 Self::ReqwestMiddlewareError(value)
682 }
683}
684
685trait ReqwestExt {
686 async fn apply_body(self, body: Body, content_type: Option<&str>) -> Self;
687}
688impl ReqwestExt for RequestBuilder {
689 async fn apply_body(self, body: Body, content_type: Option<&str>) -> Self {
690 match body {
691 Body::Bytes(bytes) => self.body(bytes),
692 Body::Form(form_data) => match content_type {
693 Some("application/x-www-form-urlencoded") => self.form(&form_data),
694 #[cfg(feature = "multipart")]
695 Some("multipart/form-data") => {
696 use blitz_traits::net::Entry;
697 use blitz_traits::net::EntryValue;
698 let mut form_data = form_data;
699 let mut form = reqwest::multipart::Form::new();
700 for Entry { name, value } in form_data.0.drain(..) {
701 form = match value {
702 EntryValue::String(value) => form.text(name, value),
703 EntryValue::File(path_buf) => form
704 .file(name, path_buf)
705 .await
706 .expect("Couldn't read form file from disk"),
707 EntryValue::EmptyFile => form.part(
708 name,
709 reqwest::multipart::Part::bytes(&[])
710 .mime_str("application/octet-stream")
711 .unwrap(),
712 ),
713 };
714 }
715 self.multipart(form)
716 }
717 _ => self,
718 },
719 Body::Empty => self,
720 }
721 }
722}
723
724struct DummyNetWaker;
725impl NetWaker for DummyNetWaker {
726 fn wake(&self, _client_id: usize) {}
727}
728
729#[cfg(test)]
730mod tests {
731 use super::*;
732 use blitz_traits::net::Url;
733
734 #[tokio::test]
741 async fn a_data_url_reports_the_mime_type_it_declares() {
742 let provider = Provider::new(None);
743 let request = Request::get(
744 Url::parse("data:application/wasm;base64,aGVsbG8=").unwrap(),
746 );
747
748 let response = provider
749 .fetch_response_async(request)
750 .await
751 .expect("a data URL resolves without a network");
752
753 assert_eq!(
754 response
755 .headers
756 .get(blitz_traits::platform::http::header::CONTENT_TYPE)
757 .and_then(|value| value.to_str().ok()),
758 Some("application/wasm"),
759 );
760 assert_eq!(response.body.as_ref(), b"hello");
761 assert_eq!(response.status, StatusCode::OK);
762 }
763
764 #[tokio::test]
769 async fn a_file_url_has_no_content_type_to_report() {
770 let path = std::env::temp_dir().join("blitz-net-fetch-response-test.txt");
771 std::fs::write(&path, b"file body").expect("a scratch file");
772
773 let provider = Provider::new(None);
774 let url = Url::from_file_path(&path).expect("an absolute path");
775 let response = provider
776 .fetch_response_async(Request::get(url))
777 .await
778 .expect("a file URL resolves without a network");
779
780 assert!(
781 response
782 .headers
783 .get(blitz_traits::platform::http::header::CONTENT_TYPE)
784 .is_none(),
785 "a file has no server to declare a type"
786 );
787 assert_eq!(response.body.as_ref(), b"file body");
788
789 let _ = std::fs::remove_file(&path);
790 }
791
792 #[tokio::test]
796 async fn fetch_async_still_returns_the_narrow_shape() {
797 let provider = Provider::new(None);
798 let (url, bytes) = provider
799 .fetch_async(Request::get(
800 Url::parse("data:text/plain;base64,aGVsbG8=").unwrap(),
801 ))
802 .await
803 .expect("a data URL resolves without a network");
804
805 assert!(url.starts_with("data:"));
806 assert_eq!(bytes.as_ref(), b"hello");
807 }
808}