#[cfg(not(target_arch = "wasm32"))]
use crate::net::cors::CorsPreflightCache;
use crate::net::cors::{self, CorsError, ResponseTainting};
use crate::net::events::NetEvent;
use crate::net::fetch_metadata::{self, RequestDestination, RequestMode, SecFetchSite};
use crate::net::fetcher_context::FetcherContext;
#[cfg(not(target_arch = "wasm32"))]
use crate::net::hsts::{self, HstsStore};
use crate::net::mixed_content::{self, MixedContentAction, MixedContentPolicy};
use crate::net::observer::NetObserver;
use crate::net::referrer::{self, ReferrerPolicy};
use crate::net::types::{BlockReason, FetchResultMeta, NetError, RequestBody, RequestCredentials};
use crate::types::PeekBuf;
use anyhow::anyhow;
use bytes::{Bytes, BytesMut};
use futures_util::{stream, StreamExt, TryStreamExt};
use http::{header, HeaderMap, Method};
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncRead, AsyncReadExt};
use tokio::time::timeout;
use tokio_util::io::StreamReader;
use tokio_util::sync::CancellationToken;
use url::{Origin, Url};
const SENSITIVE_REDIRECT_HEADERS: &[header::HeaderName] = &[
header::AUTHORIZATION,
header::COOKIE,
header::REFERER,
header::ORIGIN,
];
static REFERRER_POLICY: header::HeaderName = header::HeaderName::from_static("referrer-policy");
pub(crate) fn blocked(
observer: &Arc<dyn NetObserver + Send + Sync>,
url: Url,
reason: BlockReason,
) -> NetError {
observer.on_event(NetEvent::Blocked {
url: url.clone(),
reason,
});
NetError::Blocked { reason, url }
}
pub(crate) enum HopCheck {
Proceed(Url),
Reject(BlockReason),
}
pub(crate) fn hop_checks(
url: &Url,
mixed_content: MixedContentPolicy,
origin: Option<&Origin>,
url_allowed: &dyn Fn(&Url) -> bool,
) -> HopCheck {
if !matches!(url.scheme(), "http" | "https") {
return HopCheck::Reject(BlockReason::UnsupportedScheme);
}
let target = match mixed_content::evaluate(mixed_content, origin, url) {
MixedContentAction::Allow => url.clone(),
MixedContentAction::Upgrade(upgraded) => upgraded,
MixedContentAction::Block => return HopCheck::Reject(BlockReason::MixedContent),
};
if !url_allowed(&target) {
return HopCheck::Reject(BlockReason::UrlPolicy);
}
HopCheck::Proceed(target)
}
pub type UrlFilter = Box<dyn Fn(&Url) -> bool + Send + Sync>;
pub type CookieJarFn = Box<dyn Fn(&Url) -> Option<String> + Send + Sync>;
pub type CookieSinkFn = Box<dyn Fn(&Url, &[&str]) + Send + Sync>;
pub type ProtocolSinkFn = Box<dyn Fn(&Url, http::Version) + Send + Sync>;
pub struct NetPolicy {
pub url_allowed: UrlFilter,
pub cookies_for: CookieJarFn,
pub on_cookies: CookieSinkFn,
pub on_protocol: ProtocolSinkFn,
#[cfg(not(target_arch = "wasm32"))]
pub hsts: Option<Arc<dyn HstsStore>>,
#[cfg(not(target_arch = "wasm32"))]
pub cors_preflight: Option<Arc<dyn CorsPreflightCache>>,
}
impl Default for NetPolicy {
fn default() -> Self {
Self {
url_allowed: Box::new(|_| true),
cookies_for: Box::new(|_| None),
on_cookies: Box::new(|_, _| {}),
on_protocol: Box::new(|_, _| {}),
#[cfg(not(target_arch = "wasm32"))]
hsts: None,
#[cfg(not(target_arch = "wasm32"))]
cors_preflight: None,
}
}
}
impl NetPolicy {
pub fn from_context(ctx: &Arc<dyn FetcherContext>) -> Self {
let ctx_url = ctx.clone();
let ctx_cookies = ctx.clone();
let ctx_sink = ctx.clone();
Self {
url_allowed: Box::new(move |url| ctx_url.is_url_allowed(url)),
cookies_for: Box::new(move |url| ctx_cookies.cookies_for(url)),
on_cookies: Box::new(move |url, values| ctx_sink.on_cookies_received(url, values)),
on_protocol: Box::new(|_, _| {}),
#[cfg(not(target_arch = "wasm32"))]
hsts: None,
#[cfg(not(target_arch = "wasm32"))]
cors_preflight: None,
}
}
pub fn with_protocol_sink(mut self, sink: ProtocolSinkFn) -> Self {
self.on_protocol = sink;
self
}
#[cfg(not(target_arch = "wasm32"))]
pub fn with_hsts(mut self, store: Option<Arc<dyn HstsStore>>) -> Self {
self.hsts = store;
self
}
#[cfg(not(target_arch = "wasm32"))]
pub fn with_cors_preflight_cache(mut self, cache: Arc<dyn CorsPreflightCache>) -> Self {
self.cors_preflight = Some(cache);
self
}
#[cfg(not(target_arch = "wasm32"))]
pub fn clear_preflight_cache(mut self) -> Self {
self.cors_preflight = None;
self
}
}
pub struct RequestInit {
pub method: Method,
pub headers: HeaderMap,
pub body: Option<RequestBody>,
pub origin: Option<Origin>,
pub mixed_content: MixedContentPolicy,
pub referrer: Option<Url>,
pub referrer_policy: ReferrerPolicy,
pub destination: RequestDestination,
pub mode: RequestMode,
pub user_activated: bool,
pub credentials: RequestCredentials,
}
impl Default for RequestInit {
fn default() -> Self {
Self::get(HeaderMap::new())
}
}
impl RequestInit {
pub fn get(headers: HeaderMap) -> Self {
Self::new(Method::GET, headers, None)
}
pub fn post(headers: HeaderMap, body: impl Into<Bytes>) -> Self {
Self::new(Method::POST, headers, Some(RequestBody::bytes(body.into())))
}
pub fn new(method: Method, headers: HeaderMap, body: Option<RequestBody>) -> Self {
Self {
method,
headers,
body,
origin: None,
mixed_content: MixedContentPolicy::default(),
referrer: None,
referrer_policy: ReferrerPolicy::default(),
destination: RequestDestination::default(),
mode: RequestMode::default(),
user_activated: false,
credentials: RequestCredentials::default(),
}
}
pub fn with_referrer(mut self, referrer: Option<Url>, policy: ReferrerPolicy) -> Self {
self.referrer = referrer;
self.referrer_policy = policy;
self
}
pub fn with_fetch_metadata(
mut self,
destination: RequestDestination,
mode: RequestMode,
user_activated: bool,
) -> Self {
self.destination = destination;
self.mode = mode;
self.user_activated = user_activated;
self
}
pub fn with_mixed_content(
mut self,
origin: Option<Origin>,
policy: MixedContentPolicy,
) -> Self {
self.origin = origin;
self.mixed_content = policy;
self
}
pub fn with_credentials(mut self, credentials: RequestCredentials) -> Self {
self.credentials = credentials;
self
}
}
const PEEK_MAX: usize = 5 * 1024;
const MAX_REDIRECTS: usize = 20;
const MAX_PREALLOC: usize = 1024 * 1024;
pub struct ResponseTop {
pub meta: FetchResultMeta,
pub peek_buf: PeekBuf,
#[cfg(not(target_arch = "wasm32"))]
pub reader: Box<dyn AsyncRead + Unpin + Send>,
#[cfg(target_arch = "wasm32")]
pub reader: Box<dyn AsyncRead + Unpin>,
}
pub async fn fetch_response_top(
client: Arc<reqwest::Client>,
url: Url,
init: RequestInit,
cancel: CancellationToken,
observer: Arc<dyn NetObserver + Send + Sync>,
policy: NetPolicy,
) -> Result<ResponseTop, NetError> {
let started = Instant::now();
observer.on_event(NetEvent::Started { url: url.clone() });
let (resp, tainting) = get_with_redirects(
client.clone(),
url.clone(),
init,
cancel.clone(),
observer.clone(),
policy,
)
.await?;
let mut meta = FetchResultMeta {
tainting,
final_url: resp.url().clone(),
status: resp.status().as_u16(),
status_text: resp.status().canonical_reason().unwrap_or("").to_string(),
headers: resp.headers().clone(),
content_length: resp.content_length(), content_type: resp
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string()),
has_body: true, };
let mut body_stream = resp
.bytes_stream()
.map_err(|e| NetError::Read(Arc::new(anyhow!(e))));
let mut received_net: u64 = 0;
let mut peek_buf_vec: Vec<u8> = Vec::with_capacity(PEEK_MAX);
let mut excess: Option<Bytes> = None;
let observer_clone = observer.clone();
while peek_buf_vec.len() < PEEK_MAX {
let next = tokio::select! {
_ = cancel.cancelled() => {
observer_clone.on_event(NetEvent::Cancelled { url: url.clone(), reason: "peek stream cancelled" });
return Err(NetError::Cancelled("peek stream cancelled".into()));
}
n = body_stream.next() => n,
};
match next {
Some(Ok(chunk)) => {
received_net += chunk.len() as u64;
observer.on_event(NetEvent::Progress {
received_bytes: received_net,
elapsed: started.elapsed(),
expected_length: meta.content_length,
});
let need = PEEK_MAX.saturating_sub(peek_buf_vec.len());
if chunk.len() <= need {
peek_buf_vec.extend_from_slice(&chunk);
} else {
peek_buf_vec.extend_from_slice(&chunk[..need]);
excess = Some(chunk.slice(need..));
break;
}
}
Some(Err(e)) => {
observer.on_event(NetEvent::Failed {
url: url.clone(),
error: e.into(),
});
return Err(NetError::Read(Arc::new(anyhow!("peek read failed"))));
}
None => {
break;
}
}
}
let excess_len = excess.as_ref().map(|b| b.len() as u64).unwrap_or(0);
#[cfg(not(target_arch = "wasm32"))]
let body_stream = if let Some(ex) = excess {
stream::once(async move { Ok::<Bytes, NetError>(ex) })
.chain(body_stream)
.boxed()
} else {
body_stream.boxed()
};
#[cfg(target_arch = "wasm32")]
let body_stream = if let Some(ex) = excess {
stream::once(async move { Ok::<Bytes, NetError>(ex) })
.chain(body_stream)
.boxed_local()
} else {
body_stream.boxed_local()
};
let peek_buf = PeekBuf::from_vec(peek_buf_vec);
let has_body_by_len = meta.content_length.unwrap_or(0) > 0 || !peek_buf.is_empty();
meta.has_body = has_body_by_len;
let stream = body_stream.map_err(|e: NetError| e.to_io());
let inner_reader = StreamReader::new(stream);
let already_delivered = received_net - excess_len;
let progress_reader = ProgressReader::new(
inner_reader,
cancel.clone(),
observer.clone(),
url.clone(),
started,
meta.content_length,
already_delivered,
);
Ok(ResponseTop {
meta,
peek_buf,
reader: Box::new(progress_reader),
})
}
struct ProgressReader<R> {
inner: R,
cancel: CancellationToken,
observer: Arc<dyn NetObserver + Send + Sync>,
url: Url,
started: Instant,
expected_length: Option<u64>,
received: u64,
cancel_emitted: bool,
finished_emitted: bool,
}
impl<R: AsyncRead + Unpin> ProgressReader<R> {
fn new(
inner: R,
cancel: CancellationToken,
observer: Arc<dyn NetObserver + Send + Sync>,
url: Url,
started: Instant,
expected_length: Option<u64>,
already_received: u64,
) -> Self {
Self {
inner,
cancel,
observer,
url,
started,
expected_length,
received: already_received,
cancel_emitted: false,
finished_emitted: false,
}
}
}
impl<R: AsyncRead + Unpin> AsyncRead for ProgressReader<R> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
if self.cancel.is_cancelled() {
if !self.cancel_emitted {
self.observer.on_event(NetEvent::Cancelled {
url: self.url.clone(),
reason: "progress reader cancelled",
});
self.cancel_emitted = true;
}
let err = NetError::Cancelled("progress reader cancelled".into());
return std::task::Poll::Ready(Err(err.to_io()));
}
let pre_len = buf.filled().len();
let poll = Pin::new(&mut self.inner).poll_read(cx, buf);
if let std::task::Poll::Ready(Ok(())) = &poll {
let new_len = buf.filled().len();
let read_bytes = (new_len - pre_len) as u64;
if read_bytes == 0 && !self.finished_emitted {
self.finished_emitted = true;
self.observer.on_event(NetEvent::Finished {
received_bytes: self.received,
elapsed: self.started.elapsed(),
url: self.url.clone(),
});
}
if read_bytes > 0 {
self.received += read_bytes;
self.observer.on_event(NetEvent::Progress {
received_bytes: self.received,
elapsed: self.started.elapsed(),
expected_length: self.expected_length,
});
}
}
poll
}
}
const READ_CHUNK: usize = 16 * 1024;
#[allow(clippy::too_many_arguments)]
pub async fn fetch_response_complete(
client: Arc<reqwest::Client>,
url: Url,
init: RequestInit,
cancel: CancellationToken,
observer: Arc<dyn NetObserver + Send + Sync>,
max_bytes: Option<usize>,
read_idle_timeout: Duration,
total_body_timeout: Option<Duration>,
policy: NetPolicy,
) -> Result<(FetchResultMeta, Bytes), NetError> {
let started = Instant::now();
let ResponseTop {
meta,
peek_buf,
mut reader,
} = fetch_response_top(client, url, init, cancel.clone(), observer.clone(), policy).await?;
if let (Some(max), Some(len)) = (max_bytes, meta.content_length) {
if len as usize > max {
return Err(NetError::Read(Arc::new(anyhow!(
"content-length {} exceeds maximum size of {} bytes",
len,
max
))));
}
}
let advertised = meta.content_length.map(|n| n as usize).unwrap_or(0);
let ceiling = max_bytes.unwrap_or(MAX_PREALLOC).min(MAX_PREALLOC);
let initial_cap = advertised.min(ceiling).max(peek_buf.len());
let mut body_buf = BytesMut::with_capacity(initial_cap);
body_buf.extend_from_slice(peek_buf.as_slice());
loop {
if let Some(total) = total_body_timeout {
if started.elapsed() > total {
return Err(NetError::Timeout("total body timeout".into()));
}
}
if body_buf.capacity() - body_buf.len() < READ_CHUNK {
body_buf.reserve(READ_CHUNK);
}
let n = tokio::select! {
_ = cancel.cancelled() => {
return Err(NetError::Cancelled("fetch_request_complete cancelled".into()));
}
r = timeout(read_idle_timeout, reader.read_buf(&mut body_buf)) => {
match r {
Err(_) => return Err(NetError::Timeout("fetch_request_complete timeout".into())),
Ok(Err(e)) => return Err(NetError::Io(Arc::new(e))),
Ok(Ok(n)) => n,
}
}
};
if n == 0 {
break;
}
if let Some(max) = max_bytes {
if body_buf.len() > max {
return Err(NetError::Read(Arc::new(anyhow!(
"fetch_request_complete exceeded maximum size of {} bytes",
max
))));
}
}
}
Ok((meta, body_buf.freeze()))
}
fn send_error(
e: reqwest::Error,
url: &Url,
what: &str,
observer: &Arc<dyn NetObserver + Send + Sync>,
) -> NetError {
#[cfg(not(target_arch = "wasm32"))]
if let Some(tls) = crate::net::tls::classify(&e, url) {
observer.on_event(NetEvent::TlsFailed {
url: url.clone(),
error: tls.clone(),
});
return NetError::Tls(tls);
}
#[cfg(target_arch = "wasm32")]
let _ = (url, observer);
NetError::Read(Arc::new(anyhow::Error::from(e).context(what.to_string())))
}
async fn get_with_redirects(
client: Arc<reqwest::Client>,
url: Url,
init: RequestInit,
cancel: CancellationToken,
observer: Arc<dyn NetObserver + Send + Sync>,
policy: NetPolicy,
) -> Result<(reqwest::Response, ResponseTainting), NetError> {
let mut url = url;
let mut current_method = init.method;
let mut current_headers = init.headers;
let mut current_body = init.body;
let origin = init.origin;
let mut referrer_policy = init.referrer_policy;
let mut site = SecFetchSite::SameOrigin;
let mut origin_tainted = false;
let mut tainting = ResponseTainting::Basic;
#[cfg(not(target_arch = "wasm32"))]
let credentials_include = init.credentials == RequestCredentials::Include;
for _ in 0..MAX_REDIRECTS {
#[cfg(not(target_arch = "wasm32"))]
if let Some(ref store) = policy.hsts {
if hsts::should_upgrade(store.as_ref(), &url, chrono::Utc::now()) {
url = hsts::upgrade(&url);
}
}
match hop_checks(&url, init.mixed_content, origin.as_ref(), &|u| {
(policy.url_allowed)(u)
}) {
HopCheck::Reject(reason) => return Err(blocked(&observer, url, reason)),
HopCheck::Proceed(target) => {
if target != url {
observer.on_event(NetEvent::Warning {
url: url.clone(),
message: format!("upgraded insecure request to {target}"),
});
url = target;
}
}
}
if let Some(ref o) = origin {
let has_left_origin = origin_tainted || *o != url.origin();
if has_left_origin {
match init.mode {
RequestMode::SameOrigin => {
return Err(blocked(
&observer,
url,
BlockReason::Cors(CorsError::SameOriginMode),
));
}
RequestMode::NoCors => {
tainting = ResponseTainting::Opaque;
if !cors::is_cors_safelisted_method(¤t_method) {
return Err(blocked(
&observer,
url,
BlockReason::Cors(CorsError::UnsafeMethodForNoCors),
));
}
if !cors::unsafe_request_header_names(¤t_headers).is_empty() {
return Err(blocked(
&observer,
url,
BlockReason::Cors(CorsError::UnsafeHeaderForNoCors),
));
}
}
RequestMode::Cors => tainting = ResponseTainting::Cors,
RequestMode::Navigate | RequestMode::Websocket => {}
}
}
}
if let Some(ref source) = init.referrer {
match referrer::determine(source, referrer_policy, &url) {
Some(value) => match value.as_str().parse() {
Ok(header_value) => {
current_headers.insert(header::REFERER, header_value);
}
Err(_) => {
current_headers.remove(header::REFERER);
}
},
None => {
current_headers.remove(header::REFERER);
}
}
}
let hop_site = match origin {
Some(ref o) => {
site = site.min(fetch_metadata::classify_site(o, &url));
site
}
None => SecFetchSite::None,
};
fetch_metadata::apply_sec_fetch_headers(
&mut current_headers,
&url,
init.destination,
init.mode,
hop_site,
init.user_activated,
);
if let Some(ref o) = origin {
match fetch_metadata::origin_header_value(
o,
origin_tainted,
¤t_method,
init.mode,
referrer_policy,
&url,
)
.and_then(|v| v.parse().ok())
{
Some(value) => {
current_headers.insert(header::ORIGIN, value);
}
None => {
current_headers.remove(header::ORIGIN);
}
}
}
#[cfg(not(target_arch = "wasm32"))]
if init.mode == RequestMode::Cors && tainting == ResponseTainting::Cors {
if let Some(ref o) = origin {
let unsafe_names = cors::unsafe_request_header_names(¤t_headers);
if !cors::is_cors_safelisted_method(¤t_method) || !unsafe_names.is_empty() {
let serialized = cors::serialize_origin(o, origin_tainted);
let now = chrono::Utc::now();
let granted = policy
.cors_preflight
.as_ref()
.and_then(|c| c.get(&serialized, &url, credentials_include, now))
.is_some_and(|allows| {
allows
.permits(¤t_method, &unsafe_names, credentials_include)
.is_ok()
});
if !granted {
let mut pf_headers =
cors::preflight_request_headers(¤t_method, &unsafe_names);
if let Ok(v) = serialized.parse() {
pf_headers.insert(header::ORIGIN, v);
}
fetch_metadata::apply_sec_fetch_headers(
&mut pf_headers,
&url,
init.destination,
init.mode,
hop_site,
false,
);
observer.on_event(NetEvent::CorsPreflight { url: url.clone() });
let fut = client
.request(Method::OPTIONS, url.clone())
.headers(pf_headers)
.send();
tokio::pin!(fut);
let pf_resp = tokio::select! {
_ = cancel.cancelled() => {
observer.on_event(NetEvent::Cancelled { url: url.clone(), reason: "cancelled during CORS preflight" });
return Err(NetError::Cancelled("cancelled during CORS preflight".into()));
}
r = &mut fut => r.map_err(|e| send_error(e, &url, "CORS preflight request failed", &observer))?
};
let allows = cors::validate_preflight_response(
pf_resp.status().as_u16(),
pf_resp.headers(),
o,
origin_tainted,
credentials_include,
)
.and_then(|allows| {
allows
.permits(¤t_method, &unsafe_names, credentials_include)
.map(|()| allows)
})
.map_err(|e| blocked(&observer, url.clone(), BlockReason::Cors(e)))?;
if let Some(cache) = policy.cors_preflight.as_ref() {
cache.put(&serialized, &url, credentials_include, allows, now);
}
}
}
}
}
let attach_cookies = match init.credentials {
RequestCredentials::Include => true,
RequestCredentials::Omit => false,
RequestCredentials::SameOrigin => origin
.as_ref()
.is_none_or(|o| !origin_tainted && *o == url.origin()),
};
if attach_cookies && !current_headers.contains_key(header::COOKIE) {
if let Some(cookie_str) = (policy.cookies_for)(&url) {
if let Ok(val) = cookie_str.parse() {
current_headers.insert(header::COOKIE, val);
}
}
}
let mut req_builder = client
.request(current_method.clone(), url.clone())
.headers(current_headers.clone());
if let Some(ref body) = current_body {
let (hop_body, explicit_len) = body.to_reqwest_body()?;
if let Some(len) = explicit_len {
if !current_headers.contains_key(header::CONTENT_LENGTH) {
req_builder = req_builder.header(header::CONTENT_LENGTH, len);
}
}
req_builder = req_builder.body(hop_body);
}
let fut = req_builder.send();
tokio::pin!(fut);
let resp = tokio::select! {
_ = cancel.cancelled() => {
observer.on_event(NetEvent::Cancelled { url: url.clone(), reason: "cancelled net.get_with_redirects" });
return Err(NetError::Cancelled("cancelled net.get_with_redirects".into()));
}
r = &mut fut => r.map_err(|e| send_error(e, &url, "net.get_with_redirects request failed", &observer))?
};
#[cfg(not(target_arch = "wasm32"))]
(policy.on_protocol)(resp.url(), resp.version());
#[cfg(not(target_arch = "wasm32"))]
if let Some(ref store) = policy.hsts {
hsts::record(store.as_ref(), &url, resp.headers(), chrono::Utc::now());
}
#[cfg(not(target_arch = "wasm32"))]
if tainting == ResponseTainting::Cors {
if let Some(ref o) = origin {
if let Err(e) =
cors::cors_check(o, origin_tainted, credentials_include, resp.headers())
{
return Err(blocked(&observer, url, BlockReason::Cors(e)));
}
}
}
if !resp.status().is_redirection() {
return Ok((resp, tainting));
}
let status = resp.status().as_u16();
let from = resp.url().clone();
if let Some(updated) = resp
.headers()
.get_all(&REFERRER_POLICY)
.iter()
.filter_map(|v| v.to_str().ok())
.filter_map(ReferrerPolicy::parse_header)
.next_back()
{
referrer_policy = updated;
}
let set_cookies: Vec<&str> = resp
.headers()
.get_all(header::SET_COOKIE)
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
if !set_cookies.is_empty() {
(policy.on_cookies)(&from, &set_cookies);
current_headers.remove(header::COOKIE);
}
let loc = resp
.headers()
.get(reqwest::header::LOCATION)
.and_then(|v| v.to_str().ok())
.ok_or_else(|| {
NetError::Redirect(Arc::new(anyhow!(
"redirect status {} without Location header",
status
)))
})?;
let to = from.join(loc).map_err(|e| {
NetError::Redirect(Arc::new(anyhow!("invalid redirect URL '{}': {}", loc, e)))
})?;
if (!to.username().is_empty() || to.password().is_some())
&& (init.mode == RequestMode::Cors || from.origin() != to.origin())
{
return Err(blocked(
&observer,
to,
BlockReason::Cors(CorsError::CredentialedRedirect),
));
}
match status {
301 | 302 => {
if current_method != Method::HEAD {
current_method = Method::GET;
}
current_body = None;
current_headers.remove(header::CONTENT_TYPE);
current_headers.remove(header::CONTENT_LENGTH);
current_headers.remove(header::TRANSFER_ENCODING);
}
303 => {
current_method = Method::GET;
current_body = None;
current_headers.remove(header::CONTENT_TYPE);
current_headers.remove(header::CONTENT_LENGTH);
current_headers.remove(header::TRANSFER_ENCODING);
}
307 | 308 => {}
_ => {
if current_method != Method::HEAD {
current_method = Method::GET;
}
current_body = None;
}
}
if from.origin() != to.origin() {
for h in SENSITIVE_REDIRECT_HEADERS {
current_headers.remove(h);
}
}
if let Some(ref o) = origin {
if to.origin() != from.origin() && *o != from.origin() {
origin_tainted = true;
}
}
observer.on_event(NetEvent::Redirected {
from,
to: to.clone(),
status,
});
url = to
}
Err(NetError::Redirect(Arc::new(anyhow!("too many redirects"))))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::net::referrer::ReferrerPolicy;
use crate::net::test_support::{RecordingObserver, RouteConfig, TestServer};
use cow_utils::CowUtils;
use http::HeaderMap;
use std::sync::Mutex;
use std::time::Duration;
use tokio::io::AsyncReadExt;
use tokio_util::sync::CancellationToken;
struct TestObserver;
impl NetObserver for TestObserver {
fn on_event(&self, _: NetEvent) {}
}
fn observer() -> Arc<dyn NetObserver + Send + Sync> {
Arc::new(TestObserver)
}
fn pattern(n: usize) -> Vec<u8> {
(0..n).map(|i| (i % 251) as u8).collect()
}
fn client() -> Arc<reqwest::Client> {
Arc::new(
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.unwrap(),
)
}
async fn tls_server_and_client(
routes: Vec<(&str, RouteConfig)>,
) -> (
crate::net::test_support::TestServerHandle,
Arc<reqwest::Client>,
) {
let mut srv = TestServer::new().tls("hsts.test");
for (path, cfg) in routes {
srv = srv.route(path, cfg);
}
let srv = srv.start().await;
let cert = reqwest::Certificate::from_pem(srv.cert_pem().unwrap()).unwrap();
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.tls_certs_only([cert])
.resolve(srv.tls_domain().unwrap(), srv.socket_addr())
.build()
.unwrap();
(srv, Arc::new(client))
}
#[tokio::test(flavor = "current_thread")]
async fn untrusted_certificate_is_a_tls_error() {
let srv = TestServer::new()
.tls("tls.test")
.route("/", RouteConfig::ok(b"x".to_vec()))
.start()
.await;
let client = reqwest::Client::builder()
.use_rustls_tls()
.resolve(srv.tls_domain().unwrap(), srv.socket_addr())
.build()
.unwrap();
let rec = Arc::new(RecordingObserver::new());
let err = fetch_response_top(
Arc::new(client),
srv.url("/"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
rec.clone(),
NetPolicy::default(),
)
.await;
let tls = match err {
Ok(_) => panic!("expected an error"),
Err(NetError::Tls(tls)) => tls,
Err(other) => panic!("expected NetError::Tls, got {other:?}"),
};
assert_eq!(tls.kind, crate::net::tls::TlsErrorKind::UnknownIssuer);
assert_eq!(tls.host, "tls.test");
assert!(tls.certificate.is_none());
assert_eq!(rec.tls_errors(), vec![tls]);
}
#[tokio::test(flavor = "current_thread")]
async fn certificate_for_another_host_is_a_tls_error() {
let srv = TestServer::new()
.tls("tls.test")
.route("/", RouteConfig::ok(b"x".to_vec()))
.start()
.await;
let cert = reqwest::Certificate::from_pem(srv.cert_pem().unwrap()).unwrap();
let client = reqwest::Client::builder()
.tls_certs_only([cert])
.resolve("other.test", srv.socket_addr())
.build()
.unwrap();
let url = Url::parse(&format!("https://other.test:{}/", srv.socket_addr().port())).unwrap();
let err = fetch_response_top(
Arc::new(client),
url,
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await;
match err {
Err(NetError::Tls(tls)) => {
assert_eq!(tls.kind, crate::net::tls::TlsErrorKind::HostnameMismatch);
assert_eq!(tls.host, "other.test");
}
Err(other) => panic!("expected NetError::Tls, got {other:?}"),
Ok(_) => panic!("expected an error"),
}
}
async fn tls_error_for_validity(
not_before: crate::net::test_support::Ymd,
not_after: crate::net::test_support::Ymd,
) -> crate::net::tls::TlsError {
let srv = TestServer::new()
.tls("tls.test")
.tls_validity(not_before, not_after)
.route("/", RouteConfig::ok(b"x".to_vec()))
.start()
.await;
let cert = reqwest::Certificate::from_pem(srv.cert_pem().unwrap()).unwrap();
let client = reqwest::Client::builder()
.tls_certs_only([cert])
.resolve(srv.tls_domain().unwrap(), srv.socket_addr())
.build()
.unwrap();
match fetch_response_top(
Arc::new(client),
srv.url("/"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
{
Err(NetError::Tls(tls)) => tls,
Err(other) => panic!("expected NetError::Tls, got {other:?}"),
Ok(_) => panic!("expected an error"),
}
}
#[tokio::test(flavor = "current_thread")]
async fn expired_certificate_is_a_tls_error() {
let tls = tls_error_for_validity((2000, 1, 1), (2001, 1, 1)).await;
assert_eq!(tls.kind, crate::net::tls::TlsErrorKind::Expired, "{tls}");
}
#[tokio::test(flavor = "current_thread")]
async fn not_yet_valid_certificate_is_a_tls_error() {
let tls = tls_error_for_validity((3000, 1, 1), (3001, 1, 1)).await;
assert_eq!(
tls.kind,
crate::net::tls::TlsErrorKind::NotYetValid,
"{tls}"
);
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_is_recorded_from_a_real_https_response() {
let (srv, client) = tls_server_and_client(vec![(
"/",
RouteConfig::ok_with_headers(
&[(
"Strict-Transport-Security",
"max-age=31536000; includeSubDomains",
)],
b"hello".to_vec(),
),
)])
.await;
let store = Arc::new(crate::net::hsts::InMemoryHstsStore::new());
let res = fetch_response_top(
client,
srv.url("/"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default().with_hsts(Some(store.clone())),
)
.await;
assert!(res.is_ok(), "tls fetch failed: {:?}", res.err());
let entry = crate::net::hsts::HstsStore::load(store.as_ref(), "hsts.test")
.expect("an https response carrying the header must arm the store");
assert!(entry.include_subdomains);
assert!(!entry.is_expired(chrono::Utc::now()));
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_is_not_recorded_over_plaintext() {
let srv = TestServer::new()
.route(
"/",
RouteConfig::ok_with_headers(
&[("Strict-Transport-Security", "max-age=31536000")],
b"hello".to_vec(),
),
)
.start()
.await;
let store = Arc::new(crate::net::hsts::InMemoryHstsStore::new());
let res = fetch_response_top(
client(),
srv.url("/"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default().with_hsts(Some(store.clone())),
)
.await;
assert!(res.is_ok());
assert!(store.is_empty(), "plaintext must never arm HSTS");
}
#[tokio::test(flavor = "current_thread")]
async fn hsts_max_age_zero_disarms_over_tls() {
let (srv, client) = tls_server_and_client(vec![(
"/",
RouteConfig::ok_with_headers(
&[("Strict-Transport-Security", "max-age=0")],
b"bye".to_vec(),
),
)])
.await;
let store = Arc::new(crate::net::hsts::InMemoryHstsStore::new());
crate::net::hsts::HstsStore::store(
store.as_ref(),
"hsts.test",
crate::net::hsts::HstsEntry {
expires_at: chrono::Utc::now() + chrono::Duration::days(30),
include_subdomains: false,
},
);
let res = fetch_response_top(
client,
srv.url("/"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default().with_hsts(Some(store.clone())),
)
.await;
assert!(res.is_ok(), "tls fetch failed: {:?}", res.err());
assert!(store.is_empty(), "max-age=0 must remove the entry");
}
async fn server() -> crate::net::test_support::TestServerHandle {
let big = pattern(64 * 1024);
let big_chunks: Vec<&[u8]> = big.chunks(5_000).collect();
let exact = vec![b'Y'; super::READ_CHUNK];
TestServer::new()
.route("/big", RouteConfig::ok(vec![b'X'; 12 * 1024]))
.route("/big-chunked", RouteConfig::chunked(big_chunks))
.route("/exact-chunk", RouteConfig::chunked(vec![exact.as_slice()]))
.route("/large-cl", RouteConfig::ok(pattern(64 * 1024)))
.route("/redirect", RouteConfig::redirect_to("/big"))
.route(
"/slow",
RouteConfig::stall_mid_body(super::PEEK_MAX, Duration::from_millis(2_000)),
)
.route("/drop", RouteConfig::drop_mid_body(100, 10_000))
.route(
"/huge-cl",
RouteConfig::drop_mid_body(super::PEEK_MAX, 1 << 45),
)
.route("/xl-cl", RouteConfig::ok(pattern(2 * 1024 * 1024)))
.route(
"/login",
RouteConfig::redirect_with_cookie("/whoami", "session=abc123; Path=/"),
)
.route("/whoami", RouteConfig::echo_cookie_header())
.route("/empty", RouteConfig::ok(b""))
.route("/nohead", RouteConfig::no_location_redirect())
.route("/loop", RouteConfig::redirect_self())
.route("/hop1", RouteConfig::redirect_to("/hop2"))
.route("/hop2", RouteConfig::redirect_to("/hop3"))
.route("/hop3", RouteConfig::ok(b"final"))
.route(
"/chunked",
RouteConfig::chunked(vec![b"hel", b"lo ", b"wor", b"ld"]),
)
.start()
.await
}
#[tokio::test(flavor = "current_thread")]
async fn top_returns_peek_and_reader_rest() {
let srv = server().await;
let ResponseTop {
meta,
peek_buf,
mut reader,
} = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(peek_buf.len(), super::PEEK_MAX);
let mut rest = Vec::new();
reader.read_to_end(&mut rest).await.unwrap();
assert_eq!(peek_buf.len() + rest.len(), 12 * 1024);
assert!(meta.has_body);
assert_eq!(meta.status, 200);
}
#[tokio::test(flavor = "current_thread")]
async fn redirects_are_followed() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/redirect"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(body.len(), 12 * 1024);
assert!(meta.has_body);
}
#[tokio::test(flavor = "current_thread")]
async fn stream_body_is_uploaded_and_replayed_on_307() {
use crate::net::types::BoxedAsyncRead;
use std::sync::atomic::{AtomicUsize, Ordering};
let srv = TestServer::new()
.route("/hop", RouteConfig::redirect_307("/echo"))
.route("/echo", RouteConfig::echo_body())
.start()
.await;
const PAYLOAD: &[u8] = b"streamed payload";
let opened = Arc::new(AtomicUsize::new(0));
let counter = opened.clone();
let body = RequestBody::stream(
move || {
counter.fetch_add(1, Ordering::SeqCst);
Ok(Box::pin(PAYLOAD) as BoxedAsyncRead)
},
Some(PAYLOAD.len() as u64),
);
let (meta, echoed) = super::fetch_response_complete(
client(),
srv.url("/hop"),
RequestInit::new(Method::POST, HeaderMap::new(), Some(body)),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&echoed[..], PAYLOAD);
assert_eq!(
opened.load(Ordering::SeqCst),
2,
"307 must replay the body by opening a fresh reader"
);
}
#[tokio::test(flavor = "current_thread")]
async fn file_body_streams_from_disk() {
let srv = TestServer::new()
.route("/echo", RouteConfig::echo_body())
.start()
.await;
let tmp = tempfile::NamedTempFile::new().unwrap();
std::fs::write(tmp.path(), b"file payload").unwrap();
let body = RequestBody::file(tmp.path()).unwrap();
let (meta, echoed) = super::fetch_response_complete(
client(),
srv.url("/echo"),
RequestInit::new(Method::POST, HeaderMap::new(), Some(body)),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&echoed[..], b"file payload");
}
#[tokio::test(flavor = "current_thread")]
async fn stream_body_open_failure_fails_the_request() {
let srv = TestServer::new()
.route("/echo", RouteConfig::echo_body())
.start()
.await;
let body = RequestBody::stream(|| Err(std::io::Error::other("source is gone")), None);
let res = super::fetch_response_complete(
client(),
srv.url("/echo"),
RequestInit::new(Method::POST, HeaderMap::new(), Some(body)),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await;
assert!(
matches!(res, Err(NetError::Io(_))),
"factory failure must surface as NetError::Io, got {res:?}"
);
}
#[tokio::test(flavor = "current_thread")]
async fn idle_timeout_triggers_on_slow_body() {
let srv = server().await;
let res = super::fetch_response_complete(
client(),
srv.url("/slow"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_millis(100),
Some(Duration::from_secs(2)),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
assert!(res
.err()
.unwrap()
.to_string()
.cow_to_ascii_lowercase()
.contains("timeout"));
}
#[tokio::test(flavor = "current_thread")]
async fn cancel_during_peek_is_honored() {
let srv = server().await;
let cancel = CancellationToken::new();
let fut = super::fetch_response_top(
client(),
srv.url("/slow"),
RequestInit::get(HeaderMap::new()),
cancel.clone(),
observer(),
NetPolicy::default(),
);
cancel.cancel();
let res = fut.await;
assert!(res.is_err());
assert!(res
.err()
.unwrap()
.to_string()
.cow_to_ascii_lowercase()
.contains("cancel"));
}
#[tokio::test(flavor = "current_thread")]
async fn fetch_complete_max_bytes_exceeded() {
let srv = server().await;
let res = super::fetch_response_complete(
client(),
srv.url("/big-chunked"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
Some(100),
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
assert!(res.err().unwrap().to_string().contains("exceeded"));
}
#[tokio::test(flavor = "current_thread")]
async fn fetch_complete_cancel_mid_body() {
let srv = server().await;
let cancel = CancellationToken::new();
let fut = super::fetch_response_complete(
client(),
srv.url("/slow"),
RequestInit::get(HeaderMap::new()),
cancel.clone(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
);
cancel.cancel();
let res = fut.await;
assert!(res.is_err());
assert!(res
.err()
.unwrap()
.to_string()
.cow_to_ascii_lowercase()
.contains("cancel"));
}
#[tokio::test(flavor = "current_thread")]
async fn progress_reader_cancel_returns_error() {
let srv = server().await;
let cancel = CancellationToken::new();
let ResponseTop { mut reader, .. } = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()),
cancel.clone(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
cancel.cancel();
assert!(reader.read(&mut vec![0u8; 1024]).await.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn drop_mid_body_produces_error() {
let srv = server().await;
let res = super::fetch_response_complete(
client(),
srv.url("/drop"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn empty_body_has_no_body_flag_and_empty_peek() {
let srv = server().await;
let ResponseTop { meta, peek_buf, .. } = super::fetch_response_top(
client(),
srv.url("/empty"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert!(peek_buf.is_empty());
assert!(!meta.has_body);
}
#[tokio::test(flavor = "current_thread")]
async fn multi_hop_redirects_are_followed() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/hop1"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"final");
}
#[tokio::test(flavor = "current_thread")]
async fn cancel_during_redirect_chain() {
let srv = server().await;
let cancel = CancellationToken::new();
let fut = super::fetch_response_top(
client(),
srv.url("/hop1"),
RequestInit::get(HeaderMap::new()),
cancel.clone(),
observer(),
NetPolicy::default(),
);
cancel.cancel();
assert!(fut.await.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn chunked_body_is_assembled_correctly() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/chunked"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(3),
Some(Duration::from_secs(5)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"hello world");
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_without_location_header_errors() {
let srv = server().await;
let res = super::fetch_response_top(
client(),
srv.url("/nohead"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_loop_exceeds_max_redirects() {
let srv = server().await;
let res = super::fetch_response_top(
client(),
srv.url("/loop"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
assert!(res
.err()
.unwrap()
.to_string()
.cow_to_ascii_lowercase()
.contains("redirect"));
}
#[tokio::test(flavor = "current_thread")]
async fn url_filter_blocks_request() {
let srv = server().await;
let res = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
NetPolicy {
url_allowed: Box::new(|_| false),
..NetPolicy::default()
},
)
.await;
assert!(matches!(
res.err(),
Some(NetError::Blocked {
reason: BlockReason::UrlPolicy,
..
})
));
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_blocks_insecure_subresource() {
let res = super::fetch_response_top(
client(),
Url::parse("http://insecure.example.com/a.js").unwrap(),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Block,
),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await;
assert!(matches!(
res.err(),
Some(NetError::Blocked {
reason: BlockReason::MixedContent,
..
})
));
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_allows_loopback_subresource() {
let srv = server().await;
assert!(srv.url("/big").host_str().unwrap().contains("127.0.0.1"));
let ResponseTop { meta, .. } = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Block,
),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_ignores_insecure_initiator() {
let srv = server().await;
let ResponseTop { meta, .. } = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("http://example.com").unwrap().origin()),
MixedContentPolicy::Block,
),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_blocks_insecure_redirect_target() {
let srv = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute("http://insecure.example.com/a.js"),
)
.start()
.await;
let res = super::fetch_response_top(
client(),
srv.url("/hop"),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Block,
),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await;
match res.err() {
Some(NetError::Blocked { reason, url }) => {
assert_eq!(reason, BlockReason::MixedContent);
assert_eq!(url.as_str(), "http://insecure.example.com/a.js");
}
other => panic!("expected a mixed content block, got {other:?}"),
}
}
#[tokio::test(flavor = "current_thread")]
async fn mixed_content_upgrades_insecure_redirect_target() {
let srv = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute("http://insecure.invalid/a.js"),
)
.start()
.await;
let rec = Arc::new(RecordingObserver::new());
let res = super::fetch_response_top(
client(),
srv.url("/hop"),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Upgrade,
),
CancellationToken::new(),
rec.clone(),
NetPolicy::default(),
)
.await;
assert_eq!(
rec.warnings(),
vec!["upgraded insecure request to https://insecure.invalid/a.js"],
"the hop must be rewritten to https"
);
assert!(
!matches!(res.as_ref().err(), Some(NetError::Blocked { .. })),
"upgrade must rewrite the hop, not block it"
);
assert_eq!(rec.blocked_reason(), None);
}
async fn referer_seen_by_server(
srv: &crate::net::test_support::TestServerHandle,
path: &str,
referrer: Option<&str>,
policy: ReferrerPolicy,
) -> String {
let (_, body) = super::fetch_response_complete(
client(),
srv.url(path),
RequestInit::get(HeaderMap::new())
.with_referrer(referrer.map(|r| Url::parse(r).unwrap()), policy),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
None,
NetPolicy::default(),
)
.await
.unwrap();
String::from_utf8_lossy(&body).to_string()
}
#[tokio::test(flavor = "current_thread")]
async fn referer_header_is_sent() {
let srv = TestServer::new()
.route("/echo", RouteConfig::echo_referer_header())
.start()
.await;
assert_eq!(
referer_seen_by_server(
&srv,
"/echo",
Some("https://example.com/page?q=1#frag"),
ReferrerPolicy::default(),
)
.await,
"https://example.com/"
);
}
#[tokio::test(flavor = "current_thread")]
async fn no_referrer_sends_no_header() {
let srv = TestServer::new()
.route("/echo", RouteConfig::echo_referer_header())
.start()
.await;
assert_eq!(
referer_seen_by_server(&srv, "/echo", None, ReferrerPolicy::default()).await,
"<absent>"
);
assert_eq!(
referer_seen_by_server(
&srv,
"/echo",
Some("https://example.com/page"),
ReferrerPolicy::NoReferrer,
)
.await,
"<absent>"
);
}
#[tokio::test(flavor = "current_thread")]
async fn referer_is_recomputed_after_a_redirect() {
let home = TestServer::new()
.route("/echo", RouteConfig::echo_referer_header())
.start()
.await;
let away = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute(home.url("/echo").as_str()),
)
.route("/echo", RouteConfig::echo_referer_header())
.start()
.await;
let doc = format!("{}page?q=1", home.base_url());
let policy = ReferrerPolicy::default();
assert_eq!(
referer_seen_by_server(&away, "/echo", Some(&doc), policy).await,
home.base_url().as_str()
);
assert_eq!(
referer_seen_by_server(&away, "/hop", Some(&doc), policy).await,
doc
);
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_referrer_policy_header_applies_to_later_hops() {
let srv = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_with_referrer_policy("/echo", "no-referrer"),
)
.route("/echo", RouteConfig::echo_referer_header())
.start()
.await;
let doc = format!("{}page?q=1", srv.base_url());
let seen =
referer_seen_by_server(&srv, "/hop", Some(&doc), ReferrerPolicy::default()).await;
assert_eq!(
seen, "<absent>",
"the redirect's no-referrer policy must suppress the header on the next hop"
);
}
async fn header_seen_by_server(
srv: &crate::net::test_support::TestServerHandle,
path: &str,
init: RequestInit,
) -> String {
let (_, body) = super::fetch_response_complete(
client(),
srv.url(path),
init,
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
None,
NetPolicy::default(),
)
.await
.unwrap();
String::from_utf8_lossy(&body).to_string()
}
#[tokio::test(flavor = "current_thread")]
async fn sec_fetch_headers_are_sent_by_default() {
let srv = TestServer::new()
.route("/dest", RouteConfig::echo_request_header("sec-fetch-dest"))
.route("/mode", RouteConfig::echo_request_header("sec-fetch-mode"))
.route("/site", RouteConfig::echo_request_header("sec-fetch-site"))
.route("/user", RouteConfig::echo_request_header("sec-fetch-user"))
.start()
.await;
let cases = [
("/dest", "empty"),
("/mode", "no-cors"),
("/site", "none"),
("/user", "<absent>"),
];
for (path, expected) in cases {
assert_eq!(
header_seen_by_server(&srv, path, RequestInit::get(HeaderMap::new())).await,
expected,
"{path}"
);
}
}
#[tokio::test(flavor = "current_thread")]
async fn sec_fetch_site_reflects_the_initiating_origin() {
let srv = TestServer::new()
.route("/site", RouteConfig::echo_request_header("sec-fetch-site"))
.start()
.await;
let mut other_port = srv.base_url();
other_port.set_port(Some(1)).unwrap();
let cases = [
(srv.base_url(), "same-origin"),
(other_port, "same-site"),
(Url::parse("https://example.com").unwrap(), "cross-site"),
];
for (initiator, expected) in cases {
let init = RequestInit::get(HeaderMap::new())
.with_mixed_content(Some(initiator.origin()), MixedContentPolicy::default());
assert_eq!(
header_seen_by_server(&srv, "/site", init).await,
expected,
"{initiator}"
);
}
}
#[tokio::test(flavor = "current_thread")]
async fn sec_fetch_site_degrades_across_redirects() {
let home = TestServer::new()
.route("/site", RouteConfig::echo_request_header("sec-fetch-site"))
.start()
.await;
let away = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute(home.url("/site").as_str()),
)
.start()
.await;
let init = RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(home.base_url().origin()),
MixedContentPolicy::default(),
);
assert_eq!(
header_seen_by_server(&away, "/hop", init).await,
"same-site",
"the foreign hop must cap the value even though the final hop is same-origin"
);
}
#[tokio::test(flavor = "current_thread")]
async fn sec_fetch_user_marks_user_navigations() {
let srv = TestServer::new()
.route("/user", RouteConfig::echo_request_header("sec-fetch-user"))
.start()
.await;
let cases = [
(RequestMode::Navigate, true, "?1"),
(RequestMode::Navigate, false, "<absent>"),
(RequestMode::NoCors, true, "<absent>"),
];
for (mode, activated, expected) in cases {
let init = RequestInit::get(HeaderMap::new()).with_fetch_metadata(
RequestDestination::Document,
mode,
activated,
);
assert_eq!(
header_seen_by_server(&srv, "/user", init).await,
expected,
"{mode:?} activated={activated}"
);
}
}
#[tokio::test(flavor = "current_thread")]
async fn origin_header_is_sent_for_post_but_not_plain_get() {
let srv = TestServer::new()
.route("/origin", RouteConfig::echo_request_header("origin"))
.start()
.await;
let initiator = srv.base_url().origin();
let post = RequestInit::post(HeaderMap::new(), b"x".to_vec())
.with_mixed_content(Some(initiator.clone()), MixedContentPolicy::default());
assert_eq!(
header_seen_by_server(&srv, "/origin", post).await,
initiator.ascii_serialization()
);
let get = RequestInit::get(HeaderMap::new())
.with_mixed_content(Some(initiator), MixedContentPolicy::default());
assert_eq!(
header_seen_by_server(&srv, "/origin", get).await,
"<absent>"
);
}
#[tokio::test(flavor = "current_thread")]
async fn origin_header_becomes_null_after_a_cross_origin_redirect() {
let home = TestServer::new()
.route("/origin", RouteConfig::echo_request_header("origin"))
.start()
.await;
let away = TestServer::new()
.route(
"/hop",
RouteConfig::redirect_absolute(home.url("/origin").as_str()),
)
.start()
.await;
let init = RequestInit::get(HeaderMap::new())
.with_fetch_metadata(RequestDestination::Empty, RequestMode::Websocket, false)
.with_mixed_content(
Some(home.base_url().origin()),
MixedContentPolicy::default(),
);
assert_eq!(header_seen_by_server(&away, "/hop", init).await, "null");
}
#[tokio::test(flavor = "current_thread")]
async fn blocking_emits_a_blocked_event() {
let rec = Arc::new(RecordingObserver::new());
let res = super::fetch_response_top(
client(),
Url::parse("http://insecure.example.com/a.js").unwrap(),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Block,
),
CancellationToken::new(),
rec.clone(),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
assert_eq!(rec.blocked_reason(), Some(BlockReason::MixedContent));
}
#[tokio::test(flavor = "current_thread")]
async fn url_filter_block_emits_a_blocked_event() {
let srv = server().await;
let rec = Arc::new(RecordingObserver::new());
let res = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
rec.clone(),
NetPolicy {
url_allowed: Box::new(|_| false),
..NetPolicy::default()
},
)
.await;
assert!(res.is_err());
assert_eq!(rec.blocked_reason(), Some(BlockReason::UrlPolicy));
}
#[tokio::test(flavor = "current_thread")]
async fn url_allowlist_vets_the_upgraded_url() {
let seen: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let seen_cb = seen.clone();
let _ = super::fetch_response_top(
client(),
Url::parse("http://insecure.invalid/a.js").unwrap(),
RequestInit::get(HeaderMap::new()).with_mixed_content(
Some(Url::parse("https://example.com").unwrap().origin()),
MixedContentPolicy::Upgrade,
),
CancellationToken::new(),
observer(),
NetPolicy {
url_allowed: Box::new(move |u| {
seen_cb.lock().unwrap().push(u.to_string());
true
}),
..NetPolicy::default()
},
)
.await;
assert_eq!(
*seen.lock().unwrap(),
vec!["https://insecure.invalid/a.js"],
"the allowlist must be shown the upgraded URL, never the http original"
);
}
#[tokio::test(flavor = "current_thread")]
async fn request_headers_are_sent() {
let srv = server().await;
let mut headers = HeaderMap::new();
headers.insert(http::header::ACCEPT, "text/html".parse().unwrap());
let ResponseTop { meta, .. } = super::fetch_response_top(
client(),
srv.url("/big"),
RequestInit::get(headers),
CancellationToken::new(),
observer(),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
}
#[tokio::test(flavor = "current_thread")]
async fn large_chunked_body_without_content_length_is_assembled() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/big-chunked"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(body.len(), 64 * 1024);
assert_eq!(&body[..], pattern(64 * 1024).as_slice());
}
#[tokio::test(flavor = "current_thread")]
async fn chunked_body_exactly_read_chunk_size_is_assembled() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/exact-chunk"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(body.len(), super::READ_CHUNK);
assert!(body.iter().all(|&b| b == b'Y'));
}
#[tokio::test(flavor = "current_thread")]
async fn large_body_with_content_length_is_assembled() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/large-cl"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(meta.content_length, Some(64 * 1024));
assert_eq!(&body[..], pattern(64 * 1024).as_slice());
}
#[tokio::test(flavor = "current_thread")]
async fn max_bytes_equal_to_body_size_succeeds() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/big"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
Some(12 * 1024),
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(body.len(), 12 * 1024);
}
#[tokio::test(flavor = "current_thread")]
async fn huge_content_length_rejected_before_body_read() {
let srv = server().await;
let res = super::fetch_response_complete(
client(),
srv.url("/huge-cl"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
Some(1024),
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
let msg = res.err().unwrap().to_string();
assert!(msg.contains("content-length"), "unexpected error: {msg}");
assert!(msg.contains("exceeds"), "unexpected error: {msg}");
}
#[tokio::test(flavor = "current_thread")]
async fn huge_content_length_does_not_preallocate() {
let srv = server().await;
let res = super::fetch_response_complete(
client(),
srv.url("/huge-cl"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await;
assert!(res.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn body_larger_than_prealloc_cap_is_assembled() {
let srv = server().await;
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/xl-cl"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(meta.content_length, Some(2 * 1024 * 1024));
assert_eq!(&body[..], pattern(2 * 1024 * 1024).as_slice());
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_set_cookie_reaches_jar_and_next_hop() {
let srv = server().await;
type ReceivedCookies = Vec<(Url, Vec<String>)>;
let jar: Arc<std::sync::Mutex<Option<String>>> = Arc::new(std::sync::Mutex::new(None));
let received: Arc<std::sync::Mutex<ReceivedCookies>> =
Arc::new(std::sync::Mutex::new(Vec::new()));
let jar_read = jar.clone();
let jar_write = jar.clone();
let received_sink = received.clone();
let policy = NetPolicy {
cookies_for: Box::new(move |_| jar_read.lock().unwrap().clone()),
on_cookies: Box::new(move |url, values| {
received_sink
.lock()
.unwrap()
.push((url.clone(), values.iter().map(|v| v.to_string()).collect()));
if let Some(v) = values.first() {
let nv = v.split(';').next().unwrap_or(v).trim().to_string();
*jar_write.lock().unwrap() = Some(nv);
}
}),
..NetPolicy::default()
};
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/login"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
policy,
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"session=abc123");
let received = received.lock().unwrap();
assert_eq!(received.len(), 1);
assert_eq!(received[0].0.path(), "/login");
assert_eq!(received[0].1, vec!["session=abc123; Path=/".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_reports_protocol_of_every_hop() {
let srv = server().await;
let seen: Arc<std::sync::Mutex<Vec<(String, http::Version)>>> =
Arc::new(std::sync::Mutex::new(Vec::new()));
let sink = seen.clone();
let policy = NetPolicy::default().with_protocol_sink(Box::new(move |url, version| {
sink.lock().unwrap().push((url.path().to_string(), version));
}));
let (meta, _) = super::fetch_response_complete(
client(),
srv.url("/login"),
RequestInit::get(HeaderMap::new()),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
policy,
)
.await
.unwrap();
assert_eq!(meta.status, 200);
let seen = seen.lock().unwrap();
assert_eq!(
*seen,
vec![
("/login".to_string(), http::Version::HTTP_11),
("/whoami".to_string(), http::Version::HTTP_11),
]
);
}
#[tokio::test(flavor = "current_thread")]
async fn redirect_set_cookie_drops_stale_cookie_header() {
let srv = server().await;
let mut headers = HeaderMap::new();
headers.insert(http::header::COOKIE, "stale=1".parse().unwrap());
let (meta, body) = super::fetch_response_complete(
client(),
srv.url("/login"),
RequestInit::get(headers),
CancellationToken::new(),
observer(),
None,
Duration::from_secs(5),
Some(Duration::from_secs(10)),
NetPolicy::default(),
)
.await
.unwrap();
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"");
}
}