use crate::net::events::NetEvent;
use crate::net::fetcher_context::FetcherContext;
use crate::net::observer::NetObserver;
use crate::net::types::{FetchResultMeta, NetError};
use crate::types::PeekBuf;
use anyhow::{anyhow, Context};
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::Url;
const SENSITIVE_REDIRECT_HEADERS: &[header::HeaderName] = &[header::AUTHORIZATION, header::COOKIE];
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 struct NetPolicy {
pub url_allowed: UrlFilter,
pub cookies_for: CookieJarFn,
pub on_cookies: CookieSinkFn,
}
impl Default for NetPolicy {
fn default() -> Self {
Self {
url_allowed: Box::new(|_| true),
cookies_for: Box::new(|_| None),
on_cookies: Box::new(|_, _| {}),
}
}
}
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)),
}
}
}
pub struct RequestInit {
pub method: Method,
pub headers: HeaderMap,
pub body: Option<Bytes>,
}
impl Default for RequestInit {
fn default() -> Self {
Self::get(HeaderMap::new())
}
}
impl RequestInit {
pub fn get(headers: HeaderMap) -> Self {
Self {
method: Method::GET,
headers,
body: None,
}
}
pub fn post(headers: HeaderMap, body: impl Into<Bytes>) -> Self {
Self {
method: Method::POST,
headers,
body: Some(body.into()),
}
}
pub fn new(method: Method, headers: HeaderMap, body: Option<Bytes>) -> Self {
Self {
method,
headers,
body,
}
}
}
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,
pub reader: Box<dyn AsyncRead + Unpin + Send>,
}
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 = get_with_redirects(
client.clone(),
url.clone(),
init,
cancel.clone(),
observer.clone(),
policy,
)
.await?;
let mut meta = FetchResultMeta {
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);
let body_stream = if let Some(ex) = excess {
stream::once(async move { Ok::<Bytes, NetError>(ex) })
.chain(body_stream)
.boxed()
} else {
body_stream.boxed()
};
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()))
}
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, NetError> {
let mut url = url;
let mut current_method = init.method;
let mut current_headers = init.headers;
let mut current_body = init.body;
for _ in 0..MAX_REDIRECTS {
if !matches!(url.scheme(), "http" | "https") {
return Err(NetError::Redirect(Arc::new(anyhow!(
"unsupported URL scheme '{}': only http and https are allowed",
url.scheme()
))));
}
if !(policy.url_allowed)(&url) {
return Err(NetError::Redirect(Arc::new(anyhow!(
"URL blocked by policy: {}",
url
))));
}
if !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 {
req_builder = req_builder.body(body.clone());
}
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.context("net.get_with_redirects request failed").map_err(|e| NetError::Read(Arc::new(e)))?
};
if !resp.status().is_redirection() {
return Ok(resp);
}
let status = resp.status().as_u16();
let from = resp.url().clone();
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)))
})?;
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);
}
}
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::test_support::{RouteConfig, TestServer};
use cow_utils::CowUtils;
use http::HeaderMap;
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 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 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!(res.is_err());
assert!(res.err().unwrap().to_string().contains("blocked"));
}
#[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_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"");
}
}