Skip to main content

blitz_net/
lib.rs

1//! Networking (HTTP, filesystem, Data URIs) for Blitz
2//!
3//! Provides implementations of [`blitz_traits::net::NetProvider`], which loads
4//! the resources a document needs, and of
5//! [`blitz_traits::platform::FetchProvider`], which serves a guest's `fetch()`.
6//!
7//! **Both run on the one client.** `Provider` holds a single `reqwest::Client`,
8//! built once with HTTP/2, the cookie jar, the compression codecs and the
9//! cacache disk cache, plus one per-host semaphore enforcing the browser's cap
10//! of six concurrent requests per origin. A guest's `fetch` therefore shares
11//! the connection pool, the cookies and the cache with the page's own loads,
12//! and counts against the same concurrency budget, which is what a browser
13//! does. A second client would have given a guest its own cookie jar and its
14//! own six connections per host.
15
16use 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
35/// The `User-Agent` sent when a caller does not choose one.
36///
37/// Kept as it was so that existing consumers see no change in what servers
38/// send them. It is worth knowing that this string is a 2020 Firefox on Linux,
39/// and that a large share of the popular web serves a degraded page, a
40/// challenge stub or nothing at all to a client it does not recognise. A
41/// consumer that wants the page a current browser would get has to say so with
42/// [`Provider::with_user_agent`].
43pub const DEFAULT_USER_AGENT: &str =
44    "Mozilla/5.0 (X11; Linux x86_64; rv:60.0) Gecko/20100101 Firefox/81.0";
45
46/// Matches real browsers' per-origin cap of 6.
47const PER_HOST_MAX_CONCURRENT: usize = 6;
48
49type HostLimits = Arc<Mutex<HashMap<String, Arc<Semaphore>>>>;
50
51#[cfg(feature = "cache")]
52type Client = reqwest_middleware::ClientWithMiddleware;
53#[cfg(not(feature = "cache"))]
54type Client = reqwest::Client;
55
56#[cfg(feature = "cache")]
57type RequestBuilder = reqwest_middleware::RequestBuilder;
58#[cfg(not(feature = "cache"))]
59type RequestBuilder = reqwest::RequestBuilder;
60
61#[cfg(feature = "cache")]
62fn get_cache_path() -> std::path::PathBuf {
63    use directories::ProjectDirs;
64    let path = ProjectDirs::from("com", "DioxusLabs", "Blitz")
65        .expect("Failed to find cache directory")
66        .cache_dir()
67        .to_owned();
68    #[cfg(feature = "tracing")]
69    tracing::info!(path = ?path.display(), "Using cache dir");
70    path
71}
72
73#[cfg(target_arch = "wasm32")]
74fn spawn(fut: impl Future + 'static) {
75    wasm_bindgen_futures::spawn_local(async move {
76        fut.await;
77    });
78}
79
80#[cfg(not(target_arch = "wasm32"))]
81fn spawn<F>(fut: F)
82where
83    F: Future + Send + 'static,
84    F::Output: Send + 'static,
85{
86    tokio::spawn(fut);
87}
88
89pub struct Provider {
90    client: Client,
91    waker: Arc<dyn NetWaker>,
92    per_host_limits: HostLimits,
93    user_agent: Arc<str>,
94    #[cfg(feature = "cache")]
95    cache_manager: CACacheManager,
96}
97impl Provider {
98    pub fn new(waker: Option<Arc<dyn NetWaker>>) -> Self {
99        Self::with_user_agent(waker, DEFAULT_USER_AGENT)
100    }
101
102    /// A provider that identifies itself as `user_agent`.
103    ///
104    /// The string is sent verbatim on document, subresource and script
105    /// requests alike, because a server that decides what to serve from the
106    /// `User-Agent` will do so on every one of them, and a client that agreed
107    /// with itself on only some would fetch a page assembled for two different
108    /// browsers.
109    pub fn with_user_agent(waker: Option<Arc<dyn NetWaker>>, user_agent: &str) -> Self {
110        let builder = reqwest::Client::builder();
111        #[cfg(feature = "cookies")]
112        let builder = builder.cookie_store(true);
113        let client = builder.build().unwrap();
114
115        #[cfg(feature = "cache")]
116        let cache_manager = CACacheManager::new(get_cache_path(), true);
117
118        #[cfg(feature = "cache")]
119        let client = reqwest_middleware::ClientBuilder::new(client)
120            .with(Cache(HttpCache {
121                mode: CacheMode::Default,
122                manager: cache_manager.clone(),
123                options: HttpCacheOptions {
124                    // Evaluate cache policy as a single-user (private) cache, like a
125                    // real browser, rather than a shared/proxy cache. The default
126                    // (`shared: true`) treats any response carrying `Set-Cookie` without
127                    // an explicit `Cache-Control: public`/`immutable` as immediately
128                    // stale, forcing a revalidation request to the server on every load.
129                    // Many CDNs (e.g. Wikimedia image hosts) serve images this way, so
130                    // the shared-cache default defeats disk caching and gets us rate
131                    // limited. A private cache honours heuristic freshness instead.
132                    cache_options: Some(CacheOptions {
133                        shared: false,
134                        ..Default::default()
135                    }),
136                    ..Default::default()
137                },
138            }))
139            .build();
140
141        let waker = waker.unwrap_or(Arc::new(DummyNetWaker));
142        Self {
143            client,
144            waker,
145            per_host_limits: Arc::new(Mutex::new(HashMap::new())),
146            user_agent: Arc::from(user_agent),
147            #[cfg(feature = "cache")]
148            cache_manager,
149        }
150    }
151    pub fn shared(waker: Option<Arc<dyn NetWaker>>) -> Arc<dyn NetProvider> {
152        Arc::new(Self::new(waker))
153    }
154    pub fn shared_with_user_agent(
155        waker: Option<Arc<dyn NetWaker>>,
156        user_agent: &str,
157    ) -> Arc<dyn NetProvider> {
158        Arc::new(Self::with_user_agent(waker, user_agent))
159    }
160    /// What this provider identifies itself as.
161    pub fn user_agent(&self) -> &str {
162        &self.user_agent
163    }
164    pub fn is_empty(&self) -> bool {
165        Arc::strong_count(&self.waker) == 1
166    }
167    pub fn count(&self) -> usize {
168        Arc::strong_count(&self.waker) - 1
169    }
170
171    #[cfg(feature = "cache")]
172    pub async fn clear_cache(&self) {
173        if let Err(e) = self.cache_manager.clear().await {
174            #[cfg(feature = "tracing")]
175            tracing::error!("Failed to clear HTTP cache: {:?}", e);
176            #[cfg(not(feature = "tracing"))]
177            let _ = e;
178        }
179    }
180}
181impl Provider {
182    async fn fetch_inner(
183        client: Client,
184        request: Request,
185        per_host_limits: HostLimits,
186        user_agent: Arc<str>,
187    ) -> Result<(String, Bytes), ProviderError> {
188        match request.url.scheme() {
189            "data" => {
190                let data_url = DataUrl::process(request.url.as_str())?;
191                let decoded = data_url.decode_to_vec()?;
192                Ok((request.url.to_string(), Bytes::from(decoded.0)))
193            }
194            "file" => {
195                let file_content = std::fs::read(request.url.path())?;
196                Ok((request.url.to_string(), Bytes::from(file_content)))
197            }
198            _ => Self::fetch_http(client, request, per_host_limits, user_agent).await,
199        }
200    }
201
202    async fn fetch_http(
203        client: Client,
204        request: Request,
205        per_host_limits: HostLimits,
206        user_agent: Arc<str>,
207    ) -> Result<(String, Bytes), ProviderError> {
208        // Acquire a per-host permit, held for the duration of the request, to
209        // keep total in-flight requests per origin bounded.
210        let host_key = request
211            .url
212            .host_str()
213            .map(str::to_owned)
214            .unwrap_or_default();
215        let semaphore = {
216            let mut map = per_host_limits.lock().unwrap();
217            map.entry(host_key)
218                .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
219                .clone()
220        };
221        let _permit = semaphore
222            .acquire()
223            .await
224            .expect("per-host semaphore was closed");
225
226        let mut req = client
227            .request(request.method, request.url)
228            .headers(request.headers)
229            .header("User-Agent", &*user_agent);
230
231        if let Some(content_type) = request.content_type.as_ref() {
232            req = req.header("Content-Type", content_type);
233        }
234
235        let req = req
236            .apply_body(request.body, request.content_type.as_deref())
237            .await;
238        let response = req.send().await?;
239        let status = response.status();
240        let final_url = response.url().to_string();
241
242        if status.is_success() {
243            return Ok((final_url, response.bytes().await?));
244        }
245
246        #[cfg(feature = "tracing")]
247        tracing::warn!(
248            url = final_url.as_str(),
249            status = status.as_u16(),
250            "HTTP error status"
251        );
252        Err(ProviderError::HttpStatus {
253            status,
254            url: final_url,
255        })
256    }
257
258    #[allow(clippy::type_complexity)]
259    pub fn fetch_with_callback(
260        &self,
261        request: Request,
262        callback: Box<dyn FnOnce(Result<(String, Bytes), ProviderError>) + Send + Sync + 'static>,
263    ) {
264        #[cfg(feature = "tracing")]
265        let url = request.url.to_string();
266
267        let client = self.client.clone();
268        let per_host_limits = self.per_host_limits.clone();
269        let user_agent = self.user_agent.clone();
270        spawn(async move {
271            let result = Self::fetch_inner(client, request, per_host_limits, user_agent).await;
272
273            #[cfg(feature = "tracing")]
274            if let Err(e) = &result {
275                #[cfg(feature = "tracing")]
276                tracing::error!(url = url.as_str(), error = ?e, "Fetching");
277            } else {
278                #[cfg(feature = "tracing")]
279                tracing::info!(url = url.as_str(), "Success fetching");
280            }
281
282            callback(result);
283        });
284    }
285
286    pub async fn fetch_async(&self, request: Request) -> Result<(String, Bytes), ProviderError> {
287        #[cfg(feature = "tracing")]
288        let url = request.url.to_string();
289
290        let client = self.client.clone();
291        let per_host_limits = self.per_host_limits.clone();
292        let user_agent = self.user_agent.clone();
293        let result = Self::fetch_inner(client, request, per_host_limits, user_agent).await;
294
295        #[cfg(feature = "tracing")]
296        if let Err(e) = &result {
297            #[cfg(feature = "tracing")]
298            tracing::error!(url = url.as_str(), error = ?e, "Fetching");
299        } else {
300            #[cfg(feature = "tracing")]
301            tracing::info!(url = url.as_str(), "Success fetching");
302        }
303
304        result
305    }
306
307    /// Fetch, keeping the response metadata that [`Provider::fetch_async`]
308    /// discards.
309    ///
310    /// `fetch_async` returns `(String, Bytes)`: the final URL and the body. That
311    /// is the right shape for the overwhelmingly common case, a document or a
312    /// subresource whose bytes are the whole answer, and it stays as it is.
313    ///
314    /// What it cannot answer is what the server *said* the bytes were. An
315    /// embedder loading a WebAssembly module wants to reject
316    /// `Content-Type: text/html` before handing the bytes to a parser that will
317    /// report an offset into a file that is not a module at all. That check
318    /// needs headers, and the headers already exist: `fetch_http` reads them off
319    /// the response and drops them on the way out.
320    ///
321    /// Additive rather than a widening of `fetch_async`, deliberately. Changing
322    /// that return type touches every caller in every embedder for a need only
323    /// some of them have, and `HeaderMap` is a heap-allocated multimap the hot
324    /// path would then build and clone for every subresource. This allocates
325    /// only for the callers that ask.
326    ///
327    /// `data:` and `file:` URLs synthesise a response. A `data:` URL states its
328    /// own mime type, so that becomes a real `Content-Type`; a `file:` URL has
329    /// none, and a caller that requires one should read an absent header as
330    /// unknown rather than as a mismatch.
331    pub async fn fetch_response_async(
332        &self,
333        request: Request,
334    ) -> Result<FetchResponse, ProviderError> {
335        let url = request.url.clone();
336        match url.scheme() {
337            "data" => {
338                // Scoped so the borrow of `url` ends before it is moved into
339                // the response.
340                let (body, headers) = {
341                    let data_url = DataUrl::process(url.as_str())?;
342                    let decoded = data_url.decode_to_vec()?;
343                    let mut headers = HeaderMap::new();
344                    if let Ok(value) = data_url.mime_type().to_string().parse() {
345                        headers.insert(blitz_traits::platform::http::header::CONTENT_TYPE, value);
346                    }
347                    (Bytes::from(decoded.0), headers)
348                };
349                Ok(FetchResponse::new(url, StatusCode::OK)
350                    .headers(headers)
351                    .body(body))
352            }
353            "file" => {
354                let file_content = std::fs::read(url.path())?;
355                Ok(FetchResponse::new(url, StatusCode::OK).body(Bytes::from(file_content)))
356            }
357            _ => {
358                let client = self.client.clone();
359                let per_host_limits = self.per_host_limits.clone();
360                let user_agent = self.user_agent.clone();
361                Self::fetch_http_response(client, request, per_host_limits, user_agent).await
362            }
363        }
364    }
365
366    /// The HTTP half of [`Provider::fetch_response_async`].
367    ///
368    /// Deliberately a sibling of [`Provider::fetch_http`] rather than a wrapper
369    /// around it: that one consumes the response to get at the body and cannot
370    /// hand back what it read on the way. The per-host permit, the user agent
371    /// and the non-2xx handling are the same, so the two must be changed
372    /// together.
373    async fn fetch_http_response(
374        client: Client,
375        request: Request,
376        per_host_limits: HostLimits,
377        user_agent: Arc<str>,
378    ) -> Result<FetchResponse, ProviderError> {
379        let host_key = request
380            .url
381            .host_str()
382            .map(str::to_owned)
383            .unwrap_or_default();
384        let semaphore = {
385            let mut map = per_host_limits.lock().unwrap();
386            map.entry(host_key)
387                .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
388                .clone()
389        };
390        let _permit = semaphore
391            .acquire()
392            .await
393            .expect("per-host semaphore was closed");
394
395        let mut req = client
396            .request(request.method, request.url)
397            .headers(request.headers)
398            .header("User-Agent", &*user_agent);
399
400        if let Some(content_type) = request.content_type.as_ref() {
401            req = req.header("Content-Type", content_type);
402        }
403
404        let req = req
405            .apply_body(request.body, request.content_type.as_deref())
406            .await;
407        let response = req.send().await?;
408        let status = response.status();
409        let final_url = response.url().clone();
410
411        if !status.is_success() {
412            #[cfg(feature = "tracing")]
413            tracing::warn!(
414                url = final_url.as_str(),
415                status = status.as_u16(),
416                "HTTP error status"
417            );
418            return Err(ProviderError::HttpStatus {
419                status,
420                url: final_url.to_string(),
421            });
422        }
423
424        // Read before the body, because taking the body consumes the response.
425        let headers = response.headers().clone();
426        Ok(FetchResponse::new(final_url, status)
427            .headers(headers)
428            .body(response.bytes().await?))
429    }
430}
431
432/// The `fetch()` path.
433///
434/// Separate from [`Provider::fetch_inner`] rather than layered on it, and the
435/// reason is the whole point of the trait: `fetch_inner` returns
436/// `(String, Bytes)` and turns any non-success status into a `ProviderError`
437/// that [`NetProvider::fetch`] then logs and drops. A caller loading an image
438/// wants exactly that. A `fetch()` caller wants the 404.
439impl Provider {
440    async fn platform_fetch_inner(
441        client: Client,
442        request: FetchRequest,
443        per_host_limits: HostLimits,
444        user_agent: Arc<str>,
445    ) -> Result<FetchResponse, FetchError> {
446        match request.url.scheme() {
447            "data" => Self::platform_fetch_data(request),
448            "file" => Self::platform_fetch_file(request),
449            "http" | "https" => {
450                Self::platform_fetch_http(client, request, per_host_limits, user_agent).await
451            }
452            scheme => Err(FetchError::UnsupportedScheme(scheme.to_owned())),
453        }
454    }
455
456    /// A `data:` URL, answered as a synthetic 200.
457    ///
458    /// The status and the `Content-Type` are invented, because a data URL has
459    /// no server to supply them, and inventing them is what the fetch
460    /// specification requires: a `data:` response is a 200 whose content type
461    /// is the one encoded in the URL.
462    fn platform_fetch_data(request: FetchRequest) -> Result<FetchResponse, FetchError> {
463        let data_url = DataUrl::process(request.url.as_str())
464            .map_err(|err| FetchError::InvalidRequest(format!("{err:?}")))?;
465        let mime = data_url.mime_type().to_string();
466        let (body, _) = data_url
467            .decode_to_vec()
468            .map_err(|err| FetchError::InvalidRequest(format!("{err:?}")))?;
469
470        let mut headers = HeaderMap::new();
471        if let Ok(value) = mime.parse() {
472            // Via `reqwest`, which re-exports `http`'s header types. Naming
473            // `http` directly would mean a new dependency for one constant.
474            headers.insert(reqwest::header::CONTENT_TYPE, value);
475        }
476
477        Ok(FetchResponse::new(request.url, StatusCode::OK)
478            .headers(headers)
479            .body(Bytes::from(body)))
480    }
481
482    /// A `file:` URL, answered as a synthetic 200.
483    ///
484    /// Goes through [`Url::to_file_path`] rather than `Url::path`, which is
485    /// what [`Provider::fetch_inner`] uses. `path` hands back the URL's
486    /// percent-encoded path component as a string, so a file whose name
487    /// contains a space or a `#` is looked up under the wrong name, and on
488    /// Windows the leading slash makes it wrong outright. `to_file_path`
489    /// decodes and refuses a URL that does not name a local path.
490    ///
491    /// **This is not access control, and it is not claiming to be.** Whether a
492    /// document may read local files at all is an origin question, and origins
493    /// are not visible here; see `blitz-platform-api`, which holds the origin
494    /// and is where such a policy belongs.
495    fn platform_fetch_file(request: FetchRequest) -> Result<FetchResponse, FetchError> {
496        let path = request.url.to_file_path().map_err(|()| {
497            FetchError::InvalidRequest(format!("not a local path: {}", request.url))
498        })?;
499
500        let body = std::fs::read(path).map_err(|err| FetchError::Network(err.to_string()))?;
501
502        Ok(FetchResponse::new(request.url, StatusCode::OK).body(Bytes::from(body)))
503    }
504
505    async fn platform_fetch_http(
506        client: Client,
507        request: FetchRequest,
508        per_host_limits: HostLimits,
509        user_agent: Arc<str>,
510    ) -> Result<FetchResponse, FetchError> {
511        // The same per-origin permit the page's own loads take, so a guest
512        // cannot open more connections to a host than a browser would.
513        let host_key = request
514            .url
515            .host_str()
516            .map(str::to_owned)
517            .unwrap_or_default();
518        let semaphore = {
519            let mut map = per_host_limits.lock().unwrap();
520            map.entry(host_key)
521                .or_insert_with(|| Arc::new(Semaphore::new(PER_HOST_MAX_CONCURRENT)))
522                .clone()
523        };
524        let _permit = semaphore
525            .acquire()
526            .await
527            .expect("per-host semaphore was closed");
528
529        let mut req = client
530            .request(request.method, request.url)
531            .headers(request.headers)
532            .header("User-Agent", &*user_agent);
533
534        if let Some(body) = request.body {
535            req = req.body(body);
536        }
537
538        let response = req
539            .send()
540            .await
541            .map_err(|err| FetchError::Network(err.to_string()))?;
542
543        // Everything below here is what `fetch_inner` throws away.
544        let status = response.status();
545        let headers = response.headers().clone();
546        let url = response.url().clone();
547        let body = response
548            .bytes()
549            .await
550            .map_err(|err| FetchError::Network(err.to_string()))?;
551
552        Ok(FetchResponse::new(url, status).headers(headers).body(body))
553    }
554}
555
556impl FetchProvider for Provider {
557    fn fetch(&self, request: FetchRequest, handler: Box<dyn FetchHandler>) {
558        let client = self.client.clone();
559        let per_host_limits = self.per_host_limits.clone();
560        let user_agent = self.user_agent.clone();
561
562        #[cfg(feature = "tracing")]
563        let url = request.url.to_string();
564
565        spawn(async move {
566            let result =
567                Self::platform_fetch_inner(client, request, per_host_limits, user_agent).await;
568
569            #[cfg(feature = "tracing")]
570            match &result {
571                Ok(response) => tracing::info!(
572                    url = url.as_str(),
573                    status = response.status.as_u16(),
574                    "fetch complete"
575                ),
576                Err(error) => tracing::error!(url = url.as_str(), error = ?error, "fetch failed"),
577            }
578
579            handler.complete(result);
580        });
581    }
582}
583
584impl NetProvider for Provider {
585    fn fetch(&self, doc_id: usize, mut request: Request, handler: Box<dyn NetHandler>) {
586        let client = self.client.clone();
587        let per_host_limits = self.per_host_limits.clone();
588        let user_agent = self.user_agent.clone();
589
590        #[cfg(feature = "tracing")]
591        tracing::info!(url = request.url.as_str(), "Fetching");
592
593        let waker = self.waker.clone();
594        spawn(async move {
595            #[cfg(feature = "tracing")]
596            let url = request.url.to_string();
597
598            let signal = request.signal.take();
599            let result = if let Some(signal) = signal {
600                AbortFetch::new(
601                    signal,
602                    Box::pin(async move {
603                        Self::fetch_inner(client, request, per_host_limits, user_agent).await
604                    }),
605                )
606                .await
607            } else {
608                Self::fetch_inner(client, request, per_host_limits, user_agent).await
609            };
610
611            waker.wake(doc_id);
612
613            match result {
614                Ok((response_url, bytes)) => {
615                    handler.bytes(response_url, bytes);
616                    #[cfg(feature = "tracing")]
617                    tracing::info!(url = url.as_str(), "Success fetching");
618                }
619                Err(e) => {
620                    #[cfg(feature = "tracing")]
621                    tracing::error!(url = url.as_str(), error = ?e, "Error fetching");
622                    #[cfg(not(feature = "tracing"))]
623                    let _ = e;
624                }
625            };
626        });
627    }
628}
629
630struct AbortFetch<F, T> {
631    signal: AbortSignal,
632    future: F,
633    _rt: PhantomData<T>,
634}
635
636impl<F, T> AbortFetch<F, T> {
637    fn new(signal: AbortSignal, future: F) -> Self {
638        Self {
639            signal,
640            future,
641            _rt: PhantomData,
642        }
643    }
644}
645
646impl<F, T> Future for AbortFetch<F, T>
647where
648    F: Future + Unpin + 'static,
649    F::Output: Into<Result<T, ProviderError>> + 'static,
650    T: Unpin,
651{
652    type Output = Result<T, ProviderError>;
653
654    fn poll(
655        mut self: std::pin::Pin<&mut Self>,
656        cx: &mut std::task::Context<'_>,
657    ) -> std::task::Poll<Self::Output> {
658        if self.signal.aborted() {
659            return Poll::Ready(Err(ProviderError::Abort));
660        }
661
662        match Pin::new(&mut self.future).poll(cx) {
663            Poll::Ready(output) => Poll::Ready(output.into()),
664            Poll::Pending => Poll::Pending,
665        }
666    }
667}
668
669#[derive(Debug)]
670pub enum ProviderError {
671    Abort,
672    Io(std::io::Error),
673    DataUrl(data_url::DataUrlError),
674    DataUrlBase64(data_url::forgiving_base64::InvalidBase64),
675    ReqwestError(reqwest::Error),
676    #[cfg(feature = "cache")]
677    ReqwestMiddlewareError(reqwest_middleware::Error),
678    HttpStatus {
679        status: reqwest::StatusCode,
680        url: String,
681    },
682}
683
684impl std::fmt::Display for ProviderError {
685    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
686        match self {
687            Self::Abort => write!(f, "request aborted"),
688            Self::Io(e) => write!(f, "io error: {e}"),
689            Self::DataUrl(e) => write!(f, "data url error: {e:?}"),
690            Self::DataUrlBase64(e) => write!(f, "data url base64 error: {e:?}"),
691            Self::ReqwestError(e) => write!(f, "reqwest error: {e}"),
692            #[cfg(feature = "cache")]
693            Self::ReqwestMiddlewareError(e) => write!(f, "reqwest middleware error: {e}"),
694            Self::HttpStatus { status, url } => write!(f, "HTTP {status} for {url}"),
695        }
696    }
697}
698
699impl From<std::io::Error> for ProviderError {
700    fn from(value: std::io::Error) -> Self {
701        Self::Io(value)
702    }
703}
704
705impl From<data_url::DataUrlError> for ProviderError {
706    fn from(value: data_url::DataUrlError) -> Self {
707        Self::DataUrl(value)
708    }
709}
710
711impl From<data_url::forgiving_base64::InvalidBase64> for ProviderError {
712    fn from(value: data_url::forgiving_base64::InvalidBase64) -> Self {
713        Self::DataUrlBase64(value)
714    }
715}
716
717impl From<reqwest::Error> for ProviderError {
718    fn from(value: reqwest::Error) -> Self {
719        Self::ReqwestError(value)
720    }
721}
722
723#[cfg(feature = "cache")]
724impl From<reqwest_middleware::Error> for ProviderError {
725    fn from(value: reqwest_middleware::Error) -> Self {
726        Self::ReqwestMiddlewareError(value)
727    }
728}
729
730trait ReqwestExt {
731    async fn apply_body(self, body: Body, content_type: Option<&str>) -> Self;
732}
733impl ReqwestExt for RequestBuilder {
734    async fn apply_body(self, body: Body, content_type: Option<&str>) -> Self {
735        match body {
736            Body::Bytes(bytes) => self.body(bytes),
737            Body::Form(form_data) => match content_type {
738                Some("application/x-www-form-urlencoded") => self.form(&form_data),
739                #[cfg(feature = "multipart")]
740                Some("multipart/form-data") => {
741                    use blitz_traits::net::Entry;
742                    use blitz_traits::net::EntryValue;
743                    let mut form_data = form_data;
744                    let mut form = reqwest::multipart::Form::new();
745                    for Entry { name, value } in form_data.0.drain(..) {
746                        form = match value {
747                            EntryValue::String(value) => form.text(name, value),
748                            EntryValue::File(path_buf) => form
749                                .file(name, path_buf)
750                                .await
751                                .expect("Couldn't read form file from disk"),
752                            EntryValue::EmptyFile => form.part(
753                                name,
754                                reqwest::multipart::Part::bytes(&[])
755                                    .mime_str("application/octet-stream")
756                                    .unwrap(),
757                            ),
758                        };
759                    }
760                    self.multipart(form)
761                }
762                _ => self,
763            },
764            Body::Empty => self,
765        }
766    }
767}
768
769struct DummyNetWaker;
770impl NetWaker for DummyNetWaker {
771    fn wake(&self, _client_id: usize) {}
772}
773
774#[cfg(test)]
775mod tests {
776    use super::*;
777    use blitz_traits::net::Url;
778
779    /// A `data:` URL states its own mime type, so the synthesised response
780    /// carries a real `Content-Type` rather than nothing.
781    ///
782    /// This is the case that makes the header meaningful for a caller checking
783    /// one: a module inlined as a `data:` URL is as legitimate as a fetched
784    /// one, and refusing it for having no type would be wrong.
785    /// Serve exactly one HTTP request on loopback and hand back what was sent.
786    ///
787    /// A real socket rather than a mock, because the thing under test is a
788    /// header on the wire. Asserting that the provider stored the string it was
789    /// given would pass just as happily with the header never attached, which
790    /// is the failure this exists to catch.
791    fn capture_one_request() -> (Url, std::thread::JoinHandle<String>) {
792        use std::io::{Read, Write};
793
794        let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("loopback is available");
795        let port = listener
796            .local_addr()
797            .expect("the socket has an address")
798            .port();
799        let handle = std::thread::spawn(move || {
800            let (mut stream, _) = listener.accept().expect("the provider connects");
801            let mut seen = Vec::new();
802            let mut byte = [0u8; 1];
803            // Read only the head. Reading to EOF would block: the client keeps
804            // the connection open waiting for the response that follows.
805            while !seen.ends_with(b"\r\n\r\n") {
806                match stream.read(&mut byte) {
807                    Ok(0) | Err(_) => break,
808                    Ok(_) => seen.push(byte[0]),
809                }
810            }
811            let _ = stream.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n");
812            let _ = stream.flush();
813            String::from_utf8_lossy(&seen).to_string()
814        });
815
816        let url = Url::parse(&format!("http://127.0.0.1:{port}/")).expect("a valid loopback URL");
817        (url, handle)
818    }
819
820    /// A caller that names a user agent has it sent.
821    ///
822    /// Sites decide what markup to serve from this header, so a browser that
823    /// cannot choose one is served whatever the default is taken for, and on a
824    /// large share of the popular web that is a challenge page rather than the
825    /// document.
826    #[tokio::test]
827    async fn a_chosen_user_agent_reaches_the_server() {
828        let (url, server) = capture_one_request();
829        let provider = Provider::with_user_agent(None, "Chuzz/1.0 (a stated identity)");
830
831        let _ = provider.fetch_async(Request::get(url)).await;
832
833        let request = server
834            .join()
835            .expect("the server thread finishes")
836            .to_lowercase();
837        assert!(
838            request.contains("user-agent: chuzz/1.0 (a stated identity)"),
839            "the chosen user agent should be on the wire, got:\n{request}"
840        );
841    }
842
843    /// A caller that names none is unchanged by this being configurable.
844    #[tokio::test]
845    async fn the_default_user_agent_is_still_sent_when_none_is_chosen() {
846        let (url, server) = capture_one_request();
847        let provider = Provider::new(None);
848
849        let _ = provider.fetch_async(Request::get(url)).await;
850
851        let request = server
852            .join()
853            .expect("the server thread finishes")
854            .to_lowercase();
855        assert!(
856            request.contains(&format!(
857                "user-agent: {}",
858                DEFAULT_USER_AGENT.to_lowercase()
859            )),
860            "the default user agent should be on the wire, got:\n{request}"
861        );
862    }
863
864    #[tokio::test]
865    async fn a_data_url_reports_the_mime_type_it_declares() {
866        let provider = Provider::new(None);
867        let request = Request::get(
868            // "hello" as base64, typed as a wasm module.
869            Url::parse("data:application/wasm;base64,aGVsbG8=").unwrap(),
870        );
871
872        let response = provider
873            .fetch_response_async(request)
874            .await
875            .expect("a data URL resolves without a network");
876
877        assert_eq!(
878            response
879                .headers
880                .get(blitz_traits::platform::http::header::CONTENT_TYPE)
881                .and_then(|value| value.to_str().ok()),
882            Some("application/wasm"),
883        );
884        assert_eq!(response.body.as_ref(), b"hello");
885        assert_eq!(response.status, StatusCode::OK);
886    }
887
888    /// A `file:` URL has no server and therefore no headers. The absence has to
889    /// be an absence, not an empty string or a guess: a caller requiring a
890    /// `Content-Type` must be able to tell "nobody said" from "said the wrong
891    /// thing".
892    #[tokio::test]
893    async fn a_file_url_has_no_content_type_to_report() {
894        let path = std::env::temp_dir().join("blitz-net-fetch-response-test.txt");
895        std::fs::write(&path, b"file body").expect("a scratch file");
896
897        let provider = Provider::new(None);
898        let url = Url::from_file_path(&path).expect("an absolute path");
899        let response = provider
900            .fetch_response_async(Request::get(url))
901            .await
902            .expect("a file URL resolves without a network");
903
904        assert!(
905            response
906                .headers
907                .get(blitz_traits::platform::http::header::CONTENT_TYPE)
908                .is_none(),
909            "a file has no server to declare a type"
910        );
911        assert_eq!(response.body.as_ref(), b"file body");
912
913        let _ = std::fs::remove_file(&path);
914    }
915
916    /// `fetch_async` is untouched by this addition, which is the point of
917    /// adding a method rather than widening it: every existing caller keeps the
918    /// cheap `(String, Bytes)` shape and allocates no `HeaderMap`.
919    #[tokio::test]
920    async fn fetch_async_still_returns_the_narrow_shape() {
921        let provider = Provider::new(None);
922        let (url, bytes) = provider
923            .fetch_async(Request::get(
924                Url::parse("data:text/plain;base64,aGVsbG8=").unwrap(),
925            ))
926            .await
927            .expect("a data URL resolves without a network");
928
929        assert!(url.starts_with("data:"));
930        assert_eq!(bytes.as_ref(), b"hello");
931    }
932}