#[cfg(not(target_arch = "wasm32"))]
use crate::net::dns::DnsResolver;
use crate::net::fetch::{
blocked, fetch_response_complete, fetch_response_top, preflight, NetPolicy, Preflight,
RequestInit, ResponseTop,
};
use crate::net::fetcher_context::FetcherContext;
#[cfg(not(target_arch = "wasm32"))]
use crate::net::hsts::{self, HstsStore, InMemoryHstsStore};
use crate::net::mixed_content::MixedContentPolicy;
use crate::net::observer::NetObserver;
#[cfg(not(target_arch = "wasm32"))]
use crate::net::proxy::ProxyConfig;
use crate::net::shared_body::{ReaderOptions, SharedBody};
use crate::net::types::{FetchRequest, FetchResult, NetError, Priority};
use crate::net::utils::{short_url, spawn_named, Waiter};
use dashmap::{DashMap, Entry};
use http::header;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::{collections::VecDeque, sync::Arc, time::Duration};
use tokio::sync::{oneshot, Notify, Semaphore};
use tokio_util::sync::CancellationToken;
use url::Url;
const SHARED_MAX_CAPACITY: usize = 32;
#[derive(Clone)]
pub struct FetcherConfig {
pub global_slots: usize,
pub h1_per_origin: usize,
pub h2_per_origin: usize,
pub connect_timeout: Duration,
pub req_timeout: Duration,
pub read_idle_timeout: Duration,
pub total_body_timeout: Option<Duration>,
pub pool_max_idle_per_host: usize,
pub pool_idle_timeout: Option<Duration>,
pub tcp_keepalive: Option<Duration>,
pub user_agent: Option<String>,
#[cfg(not(target_arch = "wasm32"))]
pub hsts: Option<Arc<dyn HstsStore>>,
pub mixed_content: MixedContentPolicy,
#[cfg(not(target_arch = "wasm32"))]
pub proxy: ProxyConfig,
#[cfg(not(target_arch = "wasm32"))]
pub dns_resolver: Option<Arc<dyn DnsResolver>>,
}
impl Default for FetcherConfig {
fn default() -> Self {
Self {
global_slots: 32,
h1_per_origin: 6,
h2_per_origin: 16,
connect_timeout: Duration::from_secs(5),
req_timeout: Duration::from_secs(60),
read_idle_timeout: Duration::from_secs(15),
total_body_timeout: Some(Duration::from_secs(180)),
pool_max_idle_per_host: 6,
pool_idle_timeout: Some(Duration::from_secs(90)),
tcp_keepalive: Some(Duration::from_secs(60)),
user_agent: None,
#[cfg(not(target_arch = "wasm32"))]
hsts: Some(Arc::new(InMemoryHstsStore::new())),
mixed_content: MixedContentPolicy::default(),
#[cfg(not(target_arch = "wasm32"))]
proxy: ProxyConfig::default(),
#[cfg(not(target_arch = "wasm32"))]
dns_resolver: None,
}
}
}
pub struct FetchInflightEntry {
parent_cancel: CancellationToken,
waiter: Arc<Waiter>,
wants_streaming: AtomicBool,
subs: AtomicUsize,
done: CancellationToken,
}
impl FetchInflightEntry {
#[inline]
fn inc_sub(&self) {
self.subs.fetch_add(1, Ordering::Relaxed);
}
#[inline]
fn dec_sub_and_maybe_cancel(&self) {
if self.subs.fetch_sub(1, Ordering::AcqRel) == 1 {
self.parent_cancel.cancel();
}
}
}
struct QueueItem {
req: FetchRequest,
cancel: CancellationToken,
reply: oneshot::Sender<FetchResult>,
}
pub struct Fetcher {
client: reqwest::Client,
client_raw: reqwest::Client,
cfg: FetcherConfig,
global_slots: Arc<Semaphore>,
per_origin: Arc<DashMap<String, Arc<Semaphore>>>,
q_high: tokio::sync::Mutex<VecDeque<QueueItem>>,
q_norm: tokio::sync::Mutex<VecDeque<QueueItem>>,
q_low: tokio::sync::Mutex<VecDeque<QueueItem>>,
q_idle: tokio::sync::Mutex<VecDeque<QueueItem>>,
inflight_map: Arc<DashMap<String, Arc<FetchInflightEntry>>>,
wake: Notify,
ctx: Arc<dyn FetcherContext>,
}
impl Fetcher {
pub fn new(config: FetcherConfig, ctx: Arc<dyn FetcherContext>) -> anyhow::Result<Self> {
anyhow::ensure!(
config.global_slots > 0,
"FetcherConfig.global_slots must be >= 1"
);
anyhow::ensure!(
config.h1_per_origin > 0,
"FetcherConfig.h1_per_origin must be >= 1"
);
anyhow::ensure!(
config.h2_per_origin > 0,
"FetcherConfig.h2_per_origin must be >= 1"
);
let client = build_client(&config, true)?;
let client_raw = build_client(&config, false)?;
Ok(Self {
client,
client_raw,
cfg: config.clone(),
global_slots: Arc::new(Semaphore::new(config.global_slots)),
per_origin: Arc::new(DashMap::new()),
q_high: tokio::sync::Mutex::new(VecDeque::new()),
q_norm: tokio::sync::Mutex::new(VecDeque::new()),
q_low: tokio::sync::Mutex::new(VecDeque::new()),
q_idle: tokio::sync::Mutex::new(VecDeque::new()),
inflight_map: Arc::new(DashMap::new()),
wake: Notify::new(),
ctx,
})
}
fn origin_key(url: &Url) -> String {
url.origin().ascii_serialization()
}
fn pick_lane<'a>(
&'a self,
high: &'a mut VecDeque<QueueItem>,
norm: &'a mut VecDeque<QueueItem>,
low: &'a mut VecDeque<QueueItem>,
idle: &'a mut VecDeque<QueueItem>,
counter: &mut u8,
) -> Option<QueueItem> {
let slot = *counter as usize;
*counter = (*counter + 1) % 15;
let try_pop = |q: &mut VecDeque<QueueItem>| q.pop_front();
match slot {
0..=7 => try_pop(high)
.or_else(|| try_pop(norm))
.or_else(|| try_pop(low))
.or_else(|| try_pop(idle)),
8..=11 => try_pop(norm)
.or_else(|| try_pop(high))
.or_else(|| try_pop(low))
.or_else(|| try_pop(idle)),
12..=13 => try_pop(low)
.or_else(|| try_pop(norm))
.or_else(|| try_pop(high))
.or_else(|| try_pop(idle)),
_ => try_pop(idle)
.or_else(|| try_pop(low))
.or_else(|| try_pop(norm))
.or_else(|| try_pop(high)),
}
}
pub async fn fetch(&self, req: FetchRequest) -> FetchResult {
self.fetch_with_cancel(req, CancellationToken::new()).await
}
pub async fn fetch_with_cancel(
&self,
req: FetchRequest,
cancel: CancellationToken,
) -> FetchResult {
let (tx, rx) = oneshot::channel();
self.submit(req, cancel, tx).await;
rx.await.unwrap_or_else(|_| {
FetchResult::Error(NetError::Cancelled(
"fetcher stopped before delivering a result".into(),
))
})
}
pub async fn submit(
&self,
req: FetchRequest,
cancel: CancellationToken,
reply_tx: oneshot::Sender<FetchResult>,
) {
log::debug!("Submitting fetch request: {:?}", req);
let mut lane = match req.priority {
Priority::High => self.q_high.lock().await,
Priority::Normal => self.q_norm.lock().await,
Priority::Low => self.q_low.lock().await,
Priority::Idle => self.q_idle.lock().await,
};
lane.push_back(QueueItem {
req,
cancel,
reply: reply_tx,
});
self.wake.notify_one();
}
pub async fn run(&self, shutdown: CancellationToken) {
let mut lane_counter: u8 = 0;
loop {
if shutdown.is_cancelled() {
break;
}
let next = {
let mut high = self.q_high.lock().await;
let mut norm = self.q_norm.lock().await;
let mut low = self.q_low.lock().await;
let mut idle = self.q_idle.lock().await;
self.pick_lane(&mut high, &mut norm, &mut low, &mut idle, &mut lane_counter)
};
#[cfg_attr(target_arch = "wasm32", allow(unused_mut))]
let Some(QueueItem {
mut req,
cancel,
reply: reply_tx,
}) = next
else {
tokio::select! {
_ = self.wake.notified() => {},
_ = shutdown.cancelled() => {},
}
continue;
};
#[cfg(not(target_arch = "wasm32"))]
if let Some(ref store) = self.cfg.hsts {
if hsts::should_upgrade(store.as_ref(), &req.url, chrono::Utc::now()) {
req.url = hsts::upgrade(&req.url);
}
}
req.mixed_content = Some(effective_mixed_content(&req, &self.cfg));
let key_opt = req.generate_request_key();
let key_str = {
let base = match key_opt {
Some(k) => k,
None => format!(
"{} {} @{}",
req.method,
req.url,
chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)
),
};
format!("{};D={}", base, req.auto_decode as u8)
};
let (inflight_entry, is_leader) = match self.inflight_map.entry(key_str.clone()) {
Entry::Occupied(entry) => {
let arc = entry.get().clone();
arc.waiter.register(reply_tx, req.streaming);
arc.inc_sub();
(arc, false)
}
Entry::Vacant(v) => {
let arc = Arc::new(FetchInflightEntry {
parent_cancel: CancellationToken::new(),
waiter: Arc::new(Waiter::new()),
wants_streaming: AtomicBool::new(req.streaming),
done: CancellationToken::new(),
subs: AtomicUsize::new(0),
});
arc.waiter.register(reply_tx, req.streaming);
arc.inc_sub();
v.insert(arc.clone());
(arc, true)
}
};
if is_leader {
self.ctx.on_ref_active(req.reference);
}
let child_cancel = cancel.clone();
let entry_for_cancel = inflight_entry.clone();
let done = entry_for_cancel.done.clone();
tokio::spawn(async move {
tokio::select! {
_ = child_cancel.cancelled() => entry_for_cancel.dec_sub_and_maybe_cancel(),
_ = done.cancelled() => {}
}
});
if req.streaming {
inflight_entry
.wants_streaming
.store(true, Ordering::Relaxed);
}
if !is_leader {
continue;
}
let observer =
self.ctx
.observer_for(req.reference, req.req_id, req.kind, req.initiator);
let reject = match preflight(
&req.url,
effective_mixed_content(&req, &self.cfg),
req.origin.as_ref(),
&|u| self.ctx.is_url_allowed(u),
) {
Preflight::Reject(reason) => Some(reason),
Preflight::Proceed(_) => None,
};
if let Some(reason) = reject {
let err = FetchResult::Error(blocked(&observer, req.url.clone(), reason));
self.inflight_map.remove(&key_str);
inflight_entry.waiter.finish(err).await;
inflight_entry.done.cancel();
self.ctx.on_ref_done(req.reference);
continue;
}
let client = if req.auto_decode {
self.client.clone()
} else {
self.client_raw.clone()
};
let global = self.global_slots.clone();
let per_origin = self.per_origin.clone();
let cfg = self.cfg.clone();
let inflight = self.inflight_map.clone();
let key_for_remove = key_str.clone();
let inflight_entry2 = inflight_entry.clone();
let shutdown_child = shutdown.clone();
let req_for_task = req.clone();
let cancel_parent = inflight_entry2.parent_cancel.clone();
let ctx_clone = self.ctx.clone();
let title = format!("Fetcher: {}", short_url(&req.url, 80));
spawn_named(&title, async move {
let origin = Fetcher::origin_key(&req.url);
let slots = per_origin
.entry(origin.clone())
.or_insert_with(|| {
Arc::new(Semaphore::new(per_origin_limit_for(&cfg, &req.url)))
})
.clone();
let g = tokio::select! { p = global.acquire_owned() => Some(p), _ = shutdown_child.cancelled() => None };
if g.is_none() {
return;
}
let h = tokio::select! { p = slots.acquire_owned() => Some(p), _ = shutdown_child.cancelled() => None };
if h.is_none() {
return;
}
let should_stream =
req.streaming || inflight_entry2.wants_streaming.load(Ordering::Relaxed);
let result = if should_stream {
perform_streaming(
&client,
observer.clone(),
&req_for_task,
&cfg,
cancel_parent.clone(),
ctx_clone.clone(),
)
.await
} else {
perform_buffered(
&client,
observer.clone(),
&req_for_task,
&cfg,
cancel_parent.clone(),
ctx_clone.clone(),
)
.await
};
let fr = match &result {
Ok(fetch_result) => fetch_result.clone(),
Err(e) => FetchResult::Error(e.clone()),
};
inflight.remove(&key_for_remove);
inflight_entry2.waiter.finish(fr).await;
inflight_entry2.done.cancel();
ctx_clone.on_ref_done(req.reference);
});
}
}
}
fn effective_mixed_content(req: &FetchRequest, cfg: &FetcherConfig) -> MixedContentPolicy {
req.mixed_content.unwrap_or(cfg.mixed_content)
}
fn make_request_init(req: &FetchRequest, cfg: &FetcherConfig) -> RequestInit {
let mut headers = req.headers.clone();
let body = req.body.as_ref().map(|b| {
if let Some(ref ct) = b.content_type {
if !headers.contains_key(header::CONTENT_TYPE) {
if let Ok(val) = ct.parse() {
headers.insert(header::CONTENT_TYPE, val);
}
}
}
b.clone()
});
RequestInit::new(req.method.clone(), headers, body)
.with_mixed_content(req.origin.clone(), effective_mixed_content(req, cfg))
.with_referrer(req.referrer.clone(), req.referrer_policy)
}
fn build_client(cfg: &FetcherConfig, decode: bool) -> anyhow::Result<reqwest::Client> {
#[cfg(target_arch = "wasm32")]
{
let _ = decode;
let mut b = reqwest::Client::builder();
if let Some(ref ua) = cfg.user_agent {
b = b.user_agent(ua);
}
Ok(b.build()?)
}
#[cfg(not(target_arch = "wasm32"))]
{
let mut b = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.connection_verbose(false)
.http2_adaptive_window(true)
.connect_timeout(cfg.connect_timeout)
.timeout(cfg.req_timeout)
.pool_max_idle_per_host(cfg.pool_max_idle_per_host)
.pool_idle_timeout(cfg.pool_idle_timeout)
.tcp_keepalive(cfg.tcp_keepalive)
.use_rustls_tls()
.gzip(decode)
.brotli(decode)
.deflate(decode);
if let Some(ref ua) = cfg.user_agent {
b = b.user_agent(ua);
}
if let Some(ref resolver) = cfg.dns_resolver {
b = b.dns_resolver(Arc::new(crate::net::dns::ReqwestResolver(resolver.clone())));
}
b = cfg.proxy.apply(b)?;
Ok(b.build()?)
}
}
fn per_origin_limit_for(cfg: &FetcherConfig, url: &Url) -> usize {
match url.scheme() {
"https" => cfg.h2_per_origin,
_ => cfg.h1_per_origin,
}
}
async fn perform_streaming(
client: &reqwest::Client,
observer: Arc<dyn NetObserver + Send + Sync>,
req: &FetchRequest,
cfg: &FetcherConfig,
cancel: CancellationToken,
ctx: Arc<dyn FetcherContext>,
) -> Result<FetchResult, NetError> {
#[cfg(not(target_arch = "wasm32"))]
let policy = NetPolicy::from_context(&ctx).with_hsts(cfg.hsts.clone());
#[cfg(target_arch = "wasm32")]
let policy = NetPolicy::from_context(&ctx);
let ResponseTop {
meta,
peek_buf,
reader,
} = fetch_response_top(
Arc::new(client.clone()),
req.url.clone(),
make_request_init(req, cfg),
cancel.clone(),
observer.clone(),
policy,
)
.await?;
notify_cookies(&ctx, &meta);
let opts = ReaderOptions {
capacity: SHARED_MAX_CAPACITY,
buf_size: 16 * 1024,
cancel: Some(cancel.clone()),
idle_timeout: Some(cfg.read_idle_timeout),
total_timeout: cfg.total_body_timeout,
max_size: req
.max_bytes
.map(|max| max.saturating_sub(peek_buf.len()) as u64),
};
Ok(FetchResult::Stream {
meta,
peek_buf,
shared: SharedBody::from_reader(reader, opts),
})
}
async fn perform_buffered(
client: &reqwest::Client,
observer: Arc<dyn NetObserver + Send + Sync>,
req: &FetchRequest,
cfg: &FetcherConfig,
cancel: CancellationToken,
ctx: Arc<dyn FetcherContext>,
) -> Result<FetchResult, NetError> {
#[cfg(not(target_arch = "wasm32"))]
let policy = NetPolicy::from_context(&ctx).with_hsts(cfg.hsts.clone());
#[cfg(target_arch = "wasm32")]
let policy = NetPolicy::from_context(&ctx);
let (meta, body) = fetch_response_complete(
Arc::new(client.clone()),
req.url.clone(),
make_request_init(req, cfg),
cancel.clone(),
observer,
req.max_bytes,
cfg.read_idle_timeout,
cfg.total_body_timeout,
policy,
)
.await?;
notify_cookies(&ctx, &meta);
Ok(FetchResult::Buffered { meta, body })
}
fn notify_cookies(ctx: &Arc<dyn FetcherContext>, meta: &crate::net::types::FetchResultMeta) {
let values: Vec<&str> = meta
.headers
.get_all(header::SET_COOKIE)
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
if !values.is_empty() {
ctx.on_cookies_received(&meta.final_url, &values);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::net::fetcher_context::NullContext;
use crate::net::proxy::ProxyRule;
use crate::net::request_ref::RequestReference;
use crate::net::test_support::{RouteConfig, TestServer};
use crate::net::types::{BlockReason, FetchRequest, Initiator, ResourceKind};
use crate::types::RequestId;
use http::{HeaderMap, Method};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::oneshot;
use tokio_util::sync::CancellationToken;
use url::Url;
fn test_config() -> FetcherConfig {
FetcherConfig {
connect_timeout: Duration::from_secs(2),
req_timeout: Duration::from_secs(5),
read_idle_timeout: Duration::from_secs(2),
total_body_timeout: Some(Duration::from_secs(10)),
..FetcherConfig::default()
}
}
async fn start_server() -> crate::net::test_support::TestServerHandle {
TestServer::new()
.route(
"/slow",
RouteConfig::stall_mid_body(0, Duration::from_secs(30)),
)
.route(
"/coalesce",
RouteConfig::delay(Duration::from_millis(50), b"coalesced".to_vec()),
)
.route("/hang", RouteConfig::hang_after_connect())
.route("/fast", RouteConfig::ok(b"x"))
.route(
"/dribble-big",
RouteConfig::chunked_with_delay(
vec![&[b'X'; 1024][..]; 12],
Duration::from_millis(30),
),
)
.route(
"/timed",
RouteConfig::delay(Duration::from_millis(60), b"ok".to_vec()),
)
.route("/not-found", RouteConfig::status(404, b"not found"))
.route("/error", RouteConfig::status(500, b"server error"))
.start()
.await
}
fn make_req(url: Url, priority: Priority) -> (FetchRequest, CancellationToken) {
let req_id = RequestId::new();
let req = FetchRequest {
reference: RequestReference::Background(0),
req_id,
url,
method: Method::GET,
headers: HeaderMap::new(),
priority,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: false,
auto_decode: true,
max_bytes: None,
body: None,
};
(req, CancellationToken::new())
}
fn dummy_item(priority: Priority) -> QueueItem {
let url = Url::parse("http://example.com/").unwrap();
let req_id = RequestId::new();
let (tx, _rx) = oneshot::channel();
QueueItem {
req: FetchRequest {
reference: RequestReference::Background(0),
req_id,
url,
method: Method::GET,
headers: HeaderMap::new(),
priority,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: false,
auto_decode: true,
max_bytes: None,
body: None,
},
cancel: CancellationToken::new(),
reply: tx,
}
}
#[test]
fn pick_lane_empty_queues_returns_none() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
let mut counter = 0u8;
assert!(f
.pick_lane(
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut counter
)
.is_none());
assert_eq!(counter, 1);
}
#[test]
fn pick_lane_counter_wraps_at_15() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
let mut counter = 14u8;
f.pick_lane(
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut counter,
);
assert_eq!(counter, 0);
}
#[test]
fn pick_lane_high_preferred_at_slots_0_to_7() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
for slot in 0u8..8 {
let mut h = VecDeque::from([dummy_item(Priority::High)]);
let mut n = VecDeque::from([dummy_item(Priority::Normal)]);
let mut counter = slot;
let item = f
.pick_lane(
&mut h,
&mut n,
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut counter,
)
.unwrap();
assert_eq!(item.req.priority, Priority::High, "slot {slot}");
}
}
#[test]
fn pick_lane_norm_preferred_at_slots_8_to_11() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
for slot in 8u8..12 {
let mut h = VecDeque::from([dummy_item(Priority::High)]);
let mut n = VecDeque::from([dummy_item(Priority::Normal)]);
let mut counter = slot;
let item = f
.pick_lane(
&mut h,
&mut n,
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut counter,
)
.unwrap();
assert_eq!(item.req.priority, Priority::Normal, "slot {slot}");
}
}
#[test]
fn pick_lane_low_preferred_at_slots_12_to_13() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
for slot in 12u8..14 {
let mut l = VecDeque::from([dummy_item(Priority::Low)]);
let mut i = VecDeque::from([dummy_item(Priority::Idle)]);
let mut counter = slot;
let item = f
.pick_lane(
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut l,
&mut i,
&mut counter,
)
.unwrap();
assert_eq!(item.req.priority, Priority::Low, "slot {slot}");
}
}
#[test]
fn pick_lane_idle_preferred_at_slot_14() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
let mut i = VecDeque::from([dummy_item(Priority::Idle)]);
let mut counter = 14u8;
let item = f
.pick_lane(
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut i,
&mut counter,
)
.unwrap();
assert_eq!(item.req.priority, Priority::Idle);
}
#[test]
fn pick_lane_falls_back_when_preferred_lane_empty() {
let f = Fetcher::new(FetcherConfig::default(), Arc::new(NullContext)).unwrap();
let mut n = VecDeque::from([dummy_item(Priority::Normal)]);
let mut counter = 0u8;
let item = f
.pick_lane(
&mut VecDeque::new(),
&mut n,
&mut VecDeque::new(),
&mut VecDeque::new(),
&mut counter,
)
.unwrap();
assert_eq!(item.req.priority, Priority::Normal);
}
#[test]
fn inflight_entry_cancel_fires_when_last_sub_removed() {
let entry = FetchInflightEntry {
parent_cancel: CancellationToken::new(),
waiter: Arc::new(Waiter::new()),
wants_streaming: AtomicBool::new(false),
subs: AtomicUsize::new(0),
done: CancellationToken::new(),
};
entry.inc_sub();
entry.inc_sub();
assert!(!entry.parent_cancel.is_cancelled());
entry.dec_sub_and_maybe_cancel();
assert!(!entry.parent_cancel.is_cancelled());
entry.dec_sub_and_maybe_cancel();
assert!(entry.parent_cancel.is_cancelled());
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_buffers_response() {
let srv = start_server().await;
let base = srv.base_url();
let shutdown = CancellationToken::new();
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let f = fetcher.clone();
tokio::spawn(async move { f.run(shutdown.clone()).await });
let (req, handle) = make_req(base, Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
match rx.await.unwrap() {
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"hello");
}
other => panic!("expected Buffered, got {:?}", other),
}
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_connection_refused_gives_error() {
let shutdown = CancellationToken::new();
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let f = fetcher.clone();
tokio::spawn(async move { f.run(shutdown.clone()).await });
let (req, handle) = make_req(Url::parse("http://127.0.0.1:1/").unwrap(), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
assert!(rx.await.unwrap().is_error());
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_cancellation_yields_error() {
let srv = start_server().await;
let base = srv.base_url();
let shutdown = CancellationToken::new();
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let f = fetcher.clone();
tokio::spawn(async move { f.run(shutdown.clone()).await });
let cancel = CancellationToken::new();
let req_id = RequestId::new();
let req = FetchRequest {
reference: RequestReference::Background(0),
req_id,
url: base.join("slow").unwrap(),
headers: HeaderMap::new(),
method: Method::GET,
priority: Priority::Normal,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: false,
auto_decode: true,
max_bytes: None,
body: None,
};
let (tx, rx) = oneshot::channel();
fetcher.submit(req, cancel.clone(), tx).await;
tokio::time::sleep(Duration::from_millis(50)).await;
cancel.cancel();
let result = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap();
assert!(result.is_error());
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_priority_high_runs_first_with_single_slot() {
let srv = start_server().await;
let base = srv.base_url();
let fetcher = Arc::new(
Fetcher::new(
FetcherConfig {
global_slots: 1,
..test_config()
},
Arc::new(NullContext),
)
.unwrap(),
);
let order = Arc::new(std::sync::Mutex::new(Vec::<&'static str>::new()));
let mut join_handles = Vec::new();
for (prio, label) in [
(Priority::Idle, "idle"),
(Priority::Low, "low"),
(Priority::Normal, "normal"),
(Priority::High, "high"),
] {
let (req, handle) = make_req(base.clone(), prio);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let order_clone = order.clone();
join_handles.push(tokio::spawn(async move {
let _ = rx.await;
order_clone.lock().unwrap().push(label);
}));
}
let f = fetcher.clone();
let shutdown = CancellationToken::new();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
for jh in join_handles {
let _ = tokio::time::timeout(Duration::from_secs(5), jh).await;
}
assert_eq!(order.lock().unwrap()[0], "high");
shutdown.cancel();
}
#[test]
fn per_origin_limit_for_uses_h2_for_https_only() {
let cfg = FetcherConfig {
h1_per_origin: 3,
h2_per_origin: 8,
..FetcherConfig::default()
};
assert_eq!(
per_origin_limit_for(&cfg, &Url::parse("http://example.com/").unwrap()),
3
);
assert_eq!(
per_origin_limit_for(&cfg, &Url::parse("https://example.com/").unwrap()),
8
);
assert_eq!(
per_origin_limit_for(&cfg, &Url::parse("ftp://example.com/").unwrap()),
3
);
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_streaming_returns_stream_result() {
let srv = start_server().await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let url = srv.base_url();
let req_id = RequestId::new();
let req = FetchRequest {
reference: RequestReference::Background(0),
req_id,
url,
method: Method::GET,
headers: HeaderMap::new(),
priority: Priority::Normal,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: true,
auto_decode: true,
max_bytes: None,
body: None,
};
let cancel = CancellationToken::new();
let (tx, rx) = oneshot::channel();
fetcher.submit(req, cancel, tx).await;
let result = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap();
match result {
FetchResult::Stream {
meta,
peek_buf,
shared,
} => {
assert_eq!(meta.status, 200);
let mut reader =
crate::net::shared_body::SharedBody::combined_reader(peek_buf, shared);
let mut body = Vec::new();
tokio::io::AsyncReadExt::read_to_end(&mut reader, &mut body)
.await
.unwrap();
assert_eq!(&body[..], b"hello");
}
other => panic!("expected Stream, got {:?}", other),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_streaming_respects_max_bytes() {
use futures_util::StreamExt;
let srv = start_server().await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (mut req, _) = make_req(srv.url("/dribble-big"), Priority::Normal);
req.streaming = true;
req.max_bytes = Some(6 * 1024);
let result = tokio::time::timeout(Duration::from_secs(5), fetcher.fetch(req))
.await
.unwrap();
match result {
FetchResult::Stream { shared, .. } => {
let mut sub = shared.subscribe_stream();
let mut saw_error = false;
while let Some(chunk) = tokio::time::timeout(Duration::from_secs(5), sub.next())
.await
.unwrap()
{
if chunk.is_err() {
saw_error = true;
break;
}
}
assert!(saw_error, "stream exceeded max_bytes without an error");
}
other => panic!("expected Stream, got {:?}", other),
}
let (mut req, _) = make_req(srv.url("/dribble-big"), Priority::Normal);
req.streaming = true;
req.max_bytes = Some(12 * 1024);
let result = tokio::time::timeout(Duration::from_secs(5), fetcher.fetch(req))
.await
.unwrap();
match result {
FetchResult::Stream {
peek_buf, shared, ..
} => {
let mut reader =
crate::net::shared_body::SharedBody::combined_reader(peek_buf, shared);
let mut body = Vec::new();
tokio::io::AsyncReadExt::read_to_end(&mut reader, &mut body)
.await
.unwrap();
assert_eq!(body.len(), 12 * 1024);
}
other => panic!("expected Stream, got {:?}", other),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_non_200_status_returned_as_buffered_not_error() {
let srv = start_server().await;
let shutdown = CancellationToken::new();
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let f = fetcher.clone();
tokio::spawn(async move { f.run(shutdown.clone()).await });
for (path, expected_status, expected_body) in [
("/not-found", 404u16, &b"not found"[..]),
("/error", 500u16, &b"server error"[..]),
] {
let (req, handle) = make_req(srv.url(path), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, expected_status, "path={path}");
assert_eq!(&body[..], expected_body, "path={path}");
}
other => panic!("expected Buffered for {path}, got {:?}", other),
}
}
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_per_origin_limit_serializes_excess_requests() {
let srv = start_server().await;
let fetcher = Arc::new(
Fetcher::new(
FetcherConfig {
h1_per_origin: 1,
global_slots: 10,
..test_config()
},
Arc::new(NullContext),
)
.unwrap(),
);
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let start = std::time::Instant::now();
let mut receivers = Vec::new();
for i in 0..3usize {
let url = Url::parse(&format!("{}timed?i={}", srv.base_url(), i)).unwrap();
let (req, handle) = make_req(url, Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
receivers.push(rx);
}
for rx in receivers {
let _ = tokio::time::timeout(Duration::from_secs(5), rx)
.await
.unwrap()
.unwrap();
}
assert!(
start.elapsed() >= Duration::from_millis(100),
"requests should be serialized, elapsed: {:?}",
start.elapsed()
);
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_fetch_convenience_delivers_result() {
let srv = start_server().await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, _) = make_req(srv.url("/fast"), Priority::Normal);
let result = tokio::time::timeout(Duration::from_secs(3), fetcher.fetch(req))
.await
.unwrap();
match result {
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"x");
}
_ => panic!("expected buffered result"),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_coalesces_duplicate_requests() {
let srv = start_server().await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let mut receivers = Vec::new();
for _ in 0..5 {
let (req, handle) = make_req(srv.url("/coalesce"), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
receivers.push(rx);
}
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
for rx in receivers {
let result = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap();
assert!(
!result.is_error(),
"every subscriber should receive a result"
);
}
assert_eq!(
srv.hit_count("/coalesce"),
1,
"coalescing must deduplicate to a single HTTP request"
);
shutdown.cancel();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn fetcher_coalescing_join_never_loses_result_under_races() {
let srv = start_server().await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
for _round in 0..20 {
let mut receivers = Vec::new();
for _ in 0..30 {
let (req, handle) = make_req(srv.url("/fast"), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
receivers.push(rx);
}
for rx in receivers {
let result = tokio::time::timeout(Duration::from_secs(5), rx)
.await
.expect("subscriber timed out waiting for result")
.expect("subscriber lost the result (waiter drained before registration)");
assert!(!result.is_error());
}
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_request_timeout_fires() {
let srv = start_server().await;
let fetcher = Arc::new(
Fetcher::new(
FetcherConfig {
req_timeout: Duration::from_millis(200),
..test_config()
},
Arc::new(NullContext),
)
.unwrap(),
);
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_req(srv.url("/hang"), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let result = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap();
assert!(
result.is_error(),
"request should fail with a timeout error"
);
shutdown.cancel();
}
async fn mixed_content_blocked(
cfg_policy: MixedContentPolicy,
req_policy: Option<MixedContentPolicy>,
) -> bool {
let fetcher = Arc::new(
Fetcher::new(
FetcherConfig {
mixed_content: cfg_policy,
..test_config()
},
Arc::new(NullContext),
)
.unwrap(),
);
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (mut req, handle) = make_req(
Url::parse("http://insecure.invalid/a.js").unwrap(),
Priority::Normal,
);
req.origin = Some(Url::parse("https://example.com").unwrap().origin());
req.mixed_content = req_policy;
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let result = tokio::time::timeout(Duration::from_secs(5), rx)
.await
.unwrap()
.unwrap();
shutdown.cancel();
matches!(
result,
FetchResult::Error(NetError::Blocked {
reason: BlockReason::MixedContent,
..
})
)
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_request_override_beats_config() {
assert!(
mixed_content_blocked(MixedContentPolicy::Block, None).await,
"no override should fall back to the fetcher-wide Block"
);
assert!(
!mixed_content_blocked(MixedContentPolicy::Block, Some(MixedContentPolicy::Allow))
.await,
"a per-request Allow must override a Block config"
);
assert!(
mixed_content_blocked(MixedContentPolicy::Allow, Some(MixedContentPolicy::Block)).await,
"a per-request Block must override an Allow config"
);
assert!(
!mixed_content_blocked(MixedContentPolicy::Allow, None).await,
"no override should fall back to the fetcher-wide Allow"
);
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_blocks_mixed_content_on_redirect_target() {
let srv = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute("http://insecure.example.com/a.js"),
)
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (mut req, handle) = make_req(srv.url("/hop"), Priority::Normal);
req.origin = Some(Url::parse("https://example.com").unwrap().origin());
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let result = tokio::time::timeout(Duration::from_secs(5), rx)
.await
.unwrap()
.unwrap();
shutdown.cancel();
assert!(
matches!(
result,
FetchResult::Error(NetError::Blocked {
reason: BlockReason::MixedContent,
..
})
),
"insecure redirect target must be blocked, got {result:?}"
);
}
fn make_post_req(
url: Url,
body: crate::net::types::RequestBody,
) -> (FetchRequest, CancellationToken) {
use http::Method;
let req_id = RequestId::new();
let req = FetchRequest {
reference: RequestReference::Background(0),
req_id,
url,
method: Method::POST,
headers: HeaderMap::new(),
priority: Priority::Normal,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: false,
auto_decode: true,
max_bytes: None,
body: Some(body),
};
(req, CancellationToken::new())
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_post_body_is_sent_and_echoed() {
let srv = TestServer::new()
.route("/echo", RouteConfig::echo_body())
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_post_req(
srv.url("/echo"),
crate::net::types::RequestBody::text("{\"x\":1}"),
);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { body, meta } => {
assert_eq!(meta.status, 200);
assert_eq!(
&body[..],
b"{\"x\":1}",
"echoed body must match the POST payload"
);
}
other => panic!("expected Buffered, got {:?}", other),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_301_downgrades_post_to_get_and_drops_body() {
let srv = TestServer::new()
.route("/post-redirect", RouteConfig::redirect_to("/landing"))
.route("/landing", RouteConfig::echo_body())
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_post_req(
srv.url("/post-redirect"),
crate::net::types::RequestBody::text("original body"),
);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, 200);
assert!(
body.is_empty(),
"body must be dropped on 301 POST→GET redirect"
);
}
other => panic!("expected Buffered, got {:?}", other),
}
shutdown.cancel();
}
struct RecordingUrlPolicy {
seen: parking_lot::Mutex<Vec<String>>,
}
impl FetcherContext for RecordingUrlPolicy {
fn observer_for(
&self,
_: RequestReference,
_: RequestId,
_: ResourceKind,
_: Initiator,
) -> Arc<dyn NetObserver + Send + Sync> {
Arc::new(crate::net::null_emitter::NullEmitter)
}
fn on_ref_active(&self, _: RequestReference) {}
fn on_ref_done(&self, _: RequestReference) {}
fn is_url_allowed(&self, url: &Url) -> bool {
self.seen.lock().push(url.path().to_string());
!url.path().contains("/blocked")
}
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_url_policy_is_applied_to_redirect_targets() {
let srv = TestServer::new()
.route("/start", RouteConfig::redirect_to("/blocked"))
.route("/blocked", RouteConfig::ok(b"SHOULD NEVER BE FETCHED"))
.start()
.await;
let ctx = Arc::new(RecordingUrlPolicy {
seen: parking_lot::Mutex::new(Vec::new()),
});
let fetcher = Arc::new(Fetcher::new(test_config(), ctx.clone()).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_req(srv.url("/start"), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let result = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap();
assert!(
ctx.seen.lock().iter().any(|p| p == "/blocked"),
"redirect target must be passed to is_url_allowed, saw: {:?}",
ctx.seen.lock()
);
assert_eq!(
srv.hit_count("/blocked"),
0,
"a blocked redirect target must never be requested"
);
assert!(
matches!(
result,
FetchResult::Error(NetError::Blocked {
reason: crate::net::types::BlockReason::UrlPolicy,
..
})
),
"blocked redirect must surface as NetError::Blocked(UrlPolicy), got {:?}",
result
);
shutdown.cancel();
}
struct RecordingPolicy {
seen: parking_lot::Mutex<Vec<String>>,
allow: bool,
}
impl RecordingPolicy {
fn new(allow: bool) -> Self {
Self {
seen: parking_lot::Mutex::new(Vec::new()),
allow,
}
}
}
impl FetcherContext for RecordingPolicy {
fn observer_for(
&self,
_: RequestReference,
_: RequestId,
_: ResourceKind,
_: Initiator,
) -> Arc<dyn NetObserver + Send + Sync> {
Arc::new(crate::net::null_emitter::NullEmitter)
}
fn on_ref_active(&self, _: RequestReference) {}
fn on_ref_done(&self, _: RequestReference) {}
fn is_url_allowed(&self, url: &Url) -> bool {
self.seen.lock().push(url.as_str().to_string());
self.allow
}
}
async fn urls_seen_for(hsts: Option<Arc<dyn HstsStore>>, request: &str) -> Vec<String> {
let ctx = Arc::new(RecordingPolicy::new(false));
let cfg = FetcherConfig {
hsts,
..test_config()
};
let fetcher = Arc::new(Fetcher::new(cfg, ctx.clone()).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_req(Url::parse(request).unwrap(), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
let _ = tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap();
shutdown.cancel();
let seen = ctx.seen.lock().clone();
seen
}
fn armed_store(host: &str, include_subdomains: bool) -> Arc<InMemoryHstsStore> {
let store = Arc::new(InMemoryHstsStore::new());
store.store(
host,
crate::net::hsts::HstsEntry {
expires_at: chrono::Utc::now() + chrono::Duration::days(1),
include_subdomains,
},
);
store
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_upgrades_url_before_key_and_policy_check() {
let seen = urls_seen_for(
Some(armed_store("hsts.example", false)),
"http://hsts.example/p",
)
.await;
assert_eq!(seen, vec!["https://hsts.example/p".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_upgrades_subdomain_of_armed_host() {
let seen = urls_seen_for(
Some(armed_store("hsts.example", true)),
"http://sub.hsts.example/p",
)
.await;
assert_eq!(seen, vec!["https://sub.hsts.example/p".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_leaves_unarmed_host_alone() {
let seen = urls_seen_for(
Some(armed_store("other.example", true)),
"http://hsts.example/p",
)
.await;
assert_eq!(seen, vec!["http://hsts.example/p".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_none_disables_upgrading() {
let seen = urls_seen_for(None, "http://hsts.example/p").await;
assert_eq!(seen, vec!["http://hsts.example/p".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
#[ignore = "requires network access to hsts.badssl.com"]
async fn hsts_live_round_trip_against_badssl() {
let store = Arc::new(InMemoryHstsStore::new());
let ctx = Arc::new(RecordingPolicy::new(true));
let cfg = FetcherConfig {
hsts: Some(store.clone()),
..FetcherConfig::default()
};
let fetcher = Arc::new(Fetcher::new(cfg, ctx.clone()).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
assert!(store.load("hsts.badssl.com").is_none());
let req =
FetchRequest::builder(Method::GET, Url::parse("https://hsts.badssl.com/").unwrap())
.build();
let res = fetcher.fetch(req).await;
assert!(
!matches!(res, FetchResult::Error(_)),
"live fetch failed: {res:?}"
);
let entry = store
.load("hsts.badssl.com")
.expect("a live HSTS response must arm the store");
assert!(entry.include_subdomains);
assert!(!entry.is_expired(chrono::Utc::now()));
ctx.seen.lock().clear();
let req =
FetchRequest::builder(Method::GET, Url::parse("http://hsts.badssl.com/").unwrap())
.build();
let _ = fetcher.fetch(req).await;
let seen = ctx.seen.lock().clone();
assert!(!seen.is_empty(), "policy should have seen the request");
assert!(
seen.iter().all(|u| u.starts_with("https://")),
"plaintext must never be requested for an armed host, saw: {seen:?}"
);
ctx.seen.lock().clear();
let req = FetchRequest::builder(
Method::GET,
Url::parse("http://sub.hsts.badssl.com/").unwrap(),
)
.build();
let _ = fetcher.fetch(req).await;
let seen = ctx.seen.lock().clone();
assert!(
seen.iter().all(|u| u.starts_with("https://")),
"subdomain of an includeSubDomains host must upgrade, saw: {seen:?}"
);
shutdown.cancel();
}
fn make_req_with_decode(
url: Url,
priority: Priority,
auto_decode: bool,
) -> (FetchRequest, CancellationToken) {
let req_id = RequestId::new();
let req = FetchRequest {
reference: RequestReference::Background(0),
req_id,
url,
method: Method::GET,
headers: HeaderMap::new(),
priority,
initiator: Initiator::Other,
kind: ResourceKind::Primary,
origin: None,
mixed_content: None,
referrer: None,
referrer_policy: Default::default(),
streaming: false,
auto_decode,
max_bytes: None,
body: None,
};
(req, CancellationToken::new())
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_auto_decode_true_decompresses_gzip() {
let srv = TestServer::new()
.route("/gz", RouteConfig::gzip_ok(b"hello compressed world"))
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, cancel) = make_req_with_decode(srv.url("/gz"), Priority::Normal, true);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, cancel, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { body, .. } => {
assert_eq!(
&body[..],
b"hello compressed world",
"auto_decode=true must yield decompressed content"
);
}
other => panic!("expected Buffered, got {:?}", other),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_auto_decode_false_returns_raw_bytes() {
let srv = TestServer::new()
.route("/gz", RouteConfig::gzip_ok(b"hello compressed world"))
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, handle) = make_req_with_decode(srv.url("/gz"), Priority::Normal, false);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, handle, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { body, .. } => {
assert_ne!(
&body[..],
b"hello compressed world",
"auto_decode=false must return raw compressed bytes"
);
assert_eq!(
&body[..2],
&[0x1f, 0x8b],
"raw bytes should start with gzip magic"
);
}
other => panic!("expected Buffered, got {:?}", other),
}
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_decode_and_raw_requests_are_not_coalesced() {
let srv = TestServer::new()
.route("/gz", RouteConfig::gzip_ok(b"data"))
.start()
.await;
let fetcher = Arc::new(Fetcher::new(test_config(), Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req_dec, handle_dec) = make_req_with_decode(srv.url("/gz"), Priority::Normal, true);
let (req_raw, handle_raw) = make_req_with_decode(srv.url("/gz"), Priority::Normal, false);
let (tx_dec, rx_dec) = oneshot::channel();
let (tx_raw, rx_raw) = oneshot::channel();
fetcher.submit(req_dec, handle_dec, tx_dec).await;
fetcher.submit(req_raw, handle_raw, tx_raw).await;
let res_dec = tokio::time::timeout(Duration::from_secs(3), rx_dec)
.await
.unwrap()
.unwrap();
let res_raw = tokio::time::timeout(Duration::from_secs(3), rx_raw)
.await
.unwrap()
.unwrap();
let body_dec = match res_dec {
FetchResult::Buffered { body, .. } => body,
o => panic!("{o:?}"),
};
let body_raw = match res_raw {
FetchResult::Buffered { body, .. } => body,
o => panic!("{o:?}"),
};
assert_eq!(&body_dec[..], b"data");
assert_eq!(&body_raw[..2], &[0x1f, 0x8b]);
assert_eq!(
srv.hit_count("/gz"),
2,
"decode and raw must not be coalesced"
);
shutdown.cancel();
}
const PROXIED_URL: &str = "http://unroutable.invalid/resource";
#[tokio::test(flavor = "current_thread")]
async fn fetcher_routes_requests_through_configured_proxy() {
let srv = TestServer::new()
.route(PROXIED_URL, RouteConfig::ok(b"served by proxy"))
.start()
.await;
let cfg = FetcherConfig {
proxy: ProxyConfig::single(srv.base_url().as_str()),
..test_config()
};
let fetcher = Arc::new(Fetcher::new(cfg, Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, cancel) = make_req(Url::parse(PROXIED_URL).unwrap(), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, cancel, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { body, .. } => assert_eq!(&body[..], b"served by proxy"),
other => panic!("expected Buffered, got {:?}", other),
}
assert_eq!(srv.hit_count(PROXIED_URL), 1, "proxy should have seen it");
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn fetcher_proxy_bypass_list_is_honoured() {
let srv = TestServer::new()
.route(PROXIED_URL, RouteConfig::ok(b"served by proxy"))
.route("/direct", RouteConfig::ok(b"served directly"))
.start()
.await;
let cfg = FetcherConfig {
proxy: ProxyConfig::Rules(vec![ProxyRule::all(srv.base_url().as_str())
.bypassing(srv.socket_addr().ip().to_string())]),
..test_config()
};
let fetcher = Arc::new(Fetcher::new(cfg, Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
let (req, cancel) = make_req(srv.url("/direct"), Priority::Normal);
let (tx, rx) = oneshot::channel();
fetcher.submit(req, cancel, tx).await;
match tokio::time::timeout(Duration::from_secs(3), rx)
.await
.unwrap()
.unwrap()
{
FetchResult::Buffered { body, .. } => assert_eq!(&body[..], b"served directly"),
other => panic!("expected Buffered, got {:?}", other),
}
assert_eq!(
srv.hit_count("/direct"),
1,
"bypassed host should be requested in origin-form, not as an absolute URI"
);
shutdown.cancel();
}
#[test]
fn fetcher_new_rejects_an_unusable_proxy_url() {
let cfg = FetcherConfig {
proxy: ProxyConfig::single("not a url"),
..test_config()
};
let err = match Fetcher::new(cfg, Arc::new(NullContext)) {
Err(e) => e,
Ok(_) => panic!("an unparseable proxy URL must not build a fetcher"),
};
assert!(
err.to_string().contains("not a url"),
"error should name the offending proxy URL, got: {err}"
);
}
}