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 = "cookies")]
36pub use reqwest::cookie::CookieStore;
37
38#[cfg(feature = "cache")]
39use http_cache_reqwest::{
40 CACacheManager, Cache, CacheMode, CacheOptions, HttpCache, HttpCacheOptions,
41};
42
43pub const DEFAULT_USER_AGENT: &str =
52 "Mozilla/5.0 (X11; Linux x86_64; rv:60.0) Gecko/20100101 Firefox/81.0";
53
54const 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 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 #[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 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 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 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 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 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 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 let headers = response.headers().clone();
461 Ok(FetchResponse::new(final_url, status)
462 .headers(headers)
463 .body(response.bytes().await?))
464 }
465}
466
467impl 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 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 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 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 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 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 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 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 #[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 #[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 #[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 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 #[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 #[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}