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