use std::fmt;
use std::io::Read;
use std::time::Duration;
use iri_string::format::ToDedicatedString;
use iri_string::types::{UriAbsoluteStr, UriReferenceStr};
use ureq::ResponseExt;
pub const DEFAULT_MAX_BODY_BYTES: usize = 8_388_608;
pub const MAX_REDIRECT_HOPS: u8 = 10;
const READ_CHUNK_BYTES: usize = 8_192;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum BodyPolicy {
#[default]
Whole,
Preview,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Truncation {
Complete,
Cut,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum RedirectPolicy {
NoFollow,
#[default]
Follow,
FollowAtMost(u8),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum FailureStage {
Resolve,
Connect,
Tls,
Send,
Headers,
Body,
EofProbe,
Redirect,
TextDecode,
Deadline,
Request,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum FailureKind {
Timeout,
Transport,
Interrupted,
InvalidUtf8,
RedirectLimit,
InvalidRequest,
Resource,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct FailureCause(FailureCauseText);
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Eq)]
enum FailureCauseText {
Sanitized(String),
Opaque(OpaqueHeader),
}
impl FailureCause {
#[must_use]
pub fn message(&self) -> &str {
match self.0 {
FailureCauseText::Sanitized(ref text) => text,
FailureCauseText::Opaque(ref header) => header.escaped(),
}
}
pub(crate) fn sanitized(text: String) -> Self {
Self(FailureCauseText::Sanitized(text))
}
fn opaque(header: OpaqueHeader) -> Self {
Self(FailureCauseText::Opaque(header))
}
}
impl fmt::Display for FailureCause {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.0 {
FailureCauseText::Sanitized(ref text) => f.write_str(text),
FailureCauseText::Opaque(ref header) => f.write_str(header.escaped()),
}
}
}
impl std::error::Error for FailureCause {}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Options {
pub timeout: Duration,
pub deadline: Option<Duration>,
user_agent: String,
headers: Vec<(String, String)>,
pub max_body_bytes: usize,
pub body_policy: BodyPolicy,
pub redirect_policy: RedirectPolicy,
}
impl Default for Options {
fn default() -> Self {
Self {
timeout: Duration::from_secs(30),
deadline: None,
user_agent: format!("lgwks-std/{}", env!("CARGO_PKG_VERSION")),
headers: Vec::new(),
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
body_policy: BodyPolicy::Whole,
redirect_policy: RedirectPolicy::default(),
}
}
}
impl Options {
#[must_use]
pub fn user_agent(mut self, user_agent: impl Into<String>) -> Self {
self.user_agent = user_agent.into();
self
}
#[must_use]
pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
let name = name.into();
if name.eq_ignore_ascii_case("Idempotency-Key") {
self.headers
.retain(|existing| !existing.0.eq_ignore_ascii_case("Idempotency-Key"));
}
self.headers.push((name, value.into()));
self
}
#[must_use]
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
#[must_use]
pub fn deadline(mut self, deadline: Duration) -> Self {
self.deadline = Some(deadline);
self
}
#[must_use]
pub fn max_body_bytes(mut self, max_body_bytes: usize) -> Self {
self.max_body_bytes = max_body_bytes;
self
}
#[must_use]
pub fn body_policy(mut self, body_policy: BodyPolicy) -> Self {
self.body_policy = body_policy;
self
}
#[must_use]
pub fn headers(&self) -> &[(String, String)] {
&self.headers
}
#[must_use]
pub fn redirect_policy(mut self, redirect_policy: RedirectPolicy) -> Self {
self.redirect_policy = redirect_policy;
self
}
#[must_use]
pub fn idempotency_key(self, key: &str) -> Self {
self.header("Idempotency-Key", key)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Response {
pub status: u16,
headers: Vec<(String, String)>,
final_target: String,
redirect_chain: Vec<String>,
header_bytes: Vec<(String, Vec<u8>)>,
body: Vec<u8>,
pub truncation: Truncation,
}
impl Response {
#[must_use]
pub fn headers(&self) -> &[(String, String)] {
&self.headers
}
pub fn header_values(&self) -> impl ExactSizeIterator<Item = (&str, &[u8])> {
self.header_bytes
.iter()
.map(|header| (header.0.as_str(), header.1.as_slice()))
}
#[must_use]
pub fn final_target(&self) -> &str {
&self.final_target
}
#[must_use]
pub fn redirect_chain(&self) -> &[String] {
&self.redirect_chain
}
#[must_use]
pub fn body(&self) -> &[u8] {
&self.body
}
pub fn text(&self) -> Result<&str, Error> {
std::str::from_utf8(&self.body).map_err(|utf8_error| Error::Failure {
stage: FailureStage::TextDecode,
kind: FailureKind::InvalidUtf8,
cause: Some(FailureCause::sanitized(format!(
"valid UTF-8 ends at byte {}",
utf8_error.valid_up_to()
))),
})
}
#[must_use]
pub fn text_lossy(&self) -> std::borrow::Cow<'_, str> {
String::from_utf8_lossy(&self.body)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum Error {
InvalidUrl,
BodyTooLarge {
limit: usize,
},
Failure {
stage: FailureStage,
kind: FailureKind,
cause: Option<FailureCause>,
},
RedirectLimitTooLarge {
requested: u8,
maximum: u8,
},
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match *self {
Self::InvalidUrl => {
write!(f, "invalid http(s) URL (absolute http(s) URI required)")
}
Self::BodyTooLarge { limit } => write!(
f,
"response body reached the declared {limit}-byte ceiling; raise \
`max_body_bytes` to read it whole, or select `BodyPolicy::Preview` to keep a \
prefix"
),
Self::Failure {
stage,
kind,
ref cause,
} => {
write!(f, "HTTP {kind:?} failure during {stage:?}")?;
if let Some(detail) = cause.as_ref() {
write!(f, ": {detail}")?;
}
Ok(())
}
Self::RedirectLimitTooLarge { requested, maximum } => write!(
f,
"redirect limit {requested} exceeds the supported maximum {maximum}"
),
}
}
}
impl std::error::Error for Error {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match *self {
Self::Failure { ref cause, .. } => cause.as_ref().map(failure_cause_source),
_ => None,
}
}
}
fn failure_cause_source(cause: &FailureCause) -> &(dyn std::error::Error + 'static) {
cause
}
pub fn validate_url(url: &str) -> Result<(), Error> {
if UriAbsoluteStr::new(url).is_err() {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
}
let Some(scheme) = url.split_once(':').map(|(scheme, _)| scheme) else {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
};
if !scheme.eq_ignore_ascii_case("http") && !scheme.eq_ignore_ascii_case("https") {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
}
let Some(rest) = scheme
.len()
.checked_add(1)
.and_then(|after| url.get(after..))
else {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
};
let Some(authority) = rest.strip_prefix("//") else {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
};
let host_end = authority.find(['/', '?', '#']).unwrap_or(authority.len());
if authority[..host_end].is_empty() {
let refusal = Err(Error::InvalidUrl);
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "validate_url: returning an error to the caller");
return refusal;
}
Ok(())
}
fn agent(options: &Options) -> ureq::Agent {
let config = ureq::Agent::config_builder()
.timeout_global(None)
.timeout_resolve(Some(options.timeout))
.timeout_connect(Some(options.timeout))
.timeout_send_request(Some(options.timeout))
.timeout_send_body(Some(options.timeout))
.timeout_recv_response(Some(options.timeout))
.timeout_recv_body(Some(options.timeout))
.max_redirects(0)
.http_status_as_error(false)
.user_agent(&options.user_agent)
.build();
ureq::Agent::new_with_config(config)
}
fn redirect_limit(options: &Options) -> Result<u32, Error> {
match options.redirect_policy {
RedirectPolicy::NoFollow => Ok(0),
RedirectPolicy::Follow => Ok(u32::from(MAX_REDIRECT_HOPS)),
RedirectPolicy::FollowAtMost(requested) if requested <= MAX_REDIRECT_HOPS => {
Ok(u32::from(requested))
}
RedirectPolicy::FollowAtMost(requested) => Err(Error::RedirectLimitTooLarge {
requested,
maximum: MAX_REDIRECT_HOPS,
}),
}
}
fn sanitized_target(uri: &str) -> String {
let Some((scheme, remainder)) = uri.split_once("://") else {
return String::from("<invalid-target>");
};
let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
let authority = remainder
.get(..authority_end)
.unwrap_or_default()
.rsplit_once('@')
.map_or_else(
|| remainder.get(..authority_end).unwrap_or_default(),
|(_, host)| host,
);
let suffix = remainder.get(authority_end..).unwrap_or_default();
let path = suffix.split(['?', '#']).next().unwrap_or_default();
let path = if path.is_empty() { "/" } else { path };
format!("{scheme}://{authority}{path}")
}
fn response_of(
mut response: ureq::http::Response<ureq::Body>,
options: &Options,
redirect_chain: Vec<String>,
) -> Result<Response, Error> {
let status = response.status().as_u16();
let header_bytes: Vec<_> = response
.headers()
.iter()
.map(|(name, value)| (name.to_string(), value.as_bytes().to_vec()))
.collect();
let headers = header_bytes
.iter()
.map(|header| {
(
header.0.clone(),
String::from_utf8_lossy(&header.1).into_owned(),
)
})
.collect();
let final_target = sanitized_target(&response.get_uri().to_string());
let (body, truncation) = read_bounded(&mut response.body_mut().as_reader(), options)?;
Ok(Response {
status,
headers,
final_target,
redirect_chain,
header_bytes,
body,
truncation,
})
}
fn read_bounded(reader: &mut impl Read, options: &Options) -> Result<(Vec<u8>, Truncation), Error> {
let ceiling = options.max_body_bytes;
let mut body = Vec::new();
let mut chunk = [0_u8; READ_CHUNK_BYTES];
let truncation = loop {
let remaining = ceiling.saturating_sub(body.len());
if remaining == 0 {
break probe_for_more(reader)?;
}
let wanted = remaining.min(READ_CHUNK_BYTES);
let read = reader
.read(&mut chunk[..wanted])
.map_err(|read_error| map_read_error(read_error, FailureStage::Body))?;
if read == 0 {
break Truncation::Complete;
}
reserve_for(&mut body, read, ceiling)?;
body.extend_from_slice(&chunk[..read]);
};
if options.body_policy == BodyPolicy::Whole && truncation == Truncation::Cut {
let refusal = Err(Error::BodyTooLarge { limit: ceiling });
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "read_bounded: returning an error to the caller");
return refusal;
}
Ok((body, truncation))
}
fn reserve_for(body: &mut Vec<u8>, incoming: usize, ceiling: usize) -> Result<(), Error> {
let needed = body.len().saturating_add(incoming);
if needed <= body.capacity() {
return Ok(());
}
let target = body
.len()
.saturating_mul(2)
.max(READ_CHUNK_BYTES)
.min(ceiling)
.max(needed);
body.try_reserve_exact(target.saturating_sub(body.len()))
.map_err(|reserve_error| {
failure(
FailureStage::Body,
FailureKind::Resource,
Some(reserve_error.to_string()),
)
})
}
fn probe_for_more(reader: &mut impl Read) -> Result<Truncation, Error> {
let mut scratch = [0_u8; 1];
let read = reader
.read(&mut scratch)
.map_err(|read_error| map_read_error(read_error, FailureStage::EofProbe))?;
Ok(if read == 0 {
Truncation::Complete
} else {
Truncation::Cut
})
}
fn timeout_stage(timeout: ureq::Timeout) -> FailureStage {
match timeout {
ureq::Timeout::Resolve => FailureStage::Resolve,
ureq::Timeout::Connect => FailureStage::Connect,
ureq::Timeout::SendRequest | ureq::Timeout::SendBody | ureq::Timeout::Await100 => {
FailureStage::Send
}
ureq::Timeout::RecvResponse => FailureStage::Headers,
ureq::Timeout::RecvBody => FailureStage::Body,
ureq::Timeout::Global | ureq::Timeout::PerCall => FailureStage::Deadline,
_ => FailureStage::Deadline,
}
}
fn failure(stage: FailureStage, kind: FailureKind, cause: Option<String>) -> Error {
Error::Failure {
stage,
kind,
cause: cause.map(FailureCause::sanitized),
}
}
fn failure_unprintable_location(header: &[u8]) -> Error {
Error::Failure {
stage: FailureStage::Redirect,
kind: FailureKind::Transport,
cause: Some(FailureCause::opaque(OpaqueHeader::new(header))),
}
}
#[derive(Clone, PartialEq, Eq)]
struct OpaqueHeader {
escaped: String,
}
impl OpaqueHeader {
fn new(header: &[u8]) -> Self {
Self {
escaped: format!("{header:?}"),
}
}
fn escaped(&self) -> &str {
&self.escaped
}
}
impl fmt::Display for OpaqueHeader {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.escaped())
}
}
impl fmt::Debug for OpaqueHeader {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(self, f)
}
}
impl std::error::Error for OpaqueHeader {}
fn map_read_error(error: std::io::Error, stage: FailureStage) -> Error {
let ureq_timeout = error
.get_ref()
.and_then(|source| source.downcast_ref::<ureq::Error>())
.and_then(|source| match source {
&ureq::Error::Timeout(timeout) => Some(timeout),
_ => None,
});
let stage = match ureq_timeout {
Some(timeout) if timeout_stage(timeout) == FailureStage::Deadline => FailureStage::Deadline,
_ => stage,
};
let timed_out = ureq_timeout.is_some() || error.kind() == std::io::ErrorKind::TimedOut;
let kind = if timed_out {
FailureKind::Timeout
} else if error.kind() == std::io::ErrorKind::Interrupted {
FailureKind::Interrupted
} else {
FailureKind::Transport
};
failure(stage, kind, Some(error.to_string()))
}
fn map_error(error: ureq::Error, is_tls: bool) -> Error {
match error {
ureq::Error::Timeout(timeout) => {
failure(timeout_stage(timeout), FailureKind::Timeout, None)
}
ureq::Error::BadUri(_) => Error::InvalidUrl,
ureq::Error::HostNotFound => failure(FailureStage::Resolve, FailureKind::Transport, None),
ureq::Error::Tls(cause) => failure(
FailureStage::Tls,
FailureKind::Transport,
Some(cause.to_owned()),
),
ureq::Error::Rustls(cause) => failure(
FailureStage::Tls,
FailureKind::Transport,
Some(cause.to_string()),
),
ureq::Error::TooManyRedirects => {
failure(FailureStage::Redirect, FailureKind::RedirectLimit, None)
}
ureq::Error::RedirectFailed => {
failure(FailureStage::Redirect, FailureKind::Transport, None)
}
ureq::Error::Io(cause) => {
let stage = match cause.kind() {
std::io::ErrorKind::ConnectionRefused
| std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::AddrNotAvailable
| std::io::ErrorKind::NetworkUnreachable => FailureStage::Connect,
std::io::ErrorKind::InvalidData | std::io::ErrorKind::UnexpectedEof if is_tls => {
FailureStage::Tls
}
_ => FailureStage::Request,
};
map_read_error(cause, stage)
}
ureq::Error::Protocol(cause) => failure(
FailureStage::Headers,
FailureKind::Transport,
Some(cause.to_string()),
),
ureq::Error::ConnectionFailed => {
failure(FailureStage::Connect, FailureKind::Transport, None)
}
ureq::Error::Http(_) => failure(FailureStage::Request, FailureKind::InvalidRequest, None),
ureq::Error::BodyExceedsLimit(_) => {
failure(FailureStage::Body, FailureKind::Transport, None)
}
_ => failure(FailureStage::Request, FailureKind::Transport, None),
}
}
fn is_https(url: &str) -> bool {
url.get(..8)
.is_some_and(|scheme| scheme.eq_ignore_ascii_case("https://"))
}
pub fn get(url: &str) -> Result<Response, Error> {
get_with(url, &Options::default())
}
pub fn get_with(url: &str, options: &Options) -> Result<Response, Error> {
exchange(url, Method::Get, options)
}
pub fn post(url: &str, content_type: &str, body: &[u8]) -> Result<Response, Error> {
post_with(url, content_type, body, &Options::default())
}
pub fn post_with(
url: &str,
content_type: &str,
body: &[u8],
options: &Options,
) -> Result<Response, Error> {
exchange(url, Method::Post { content_type, body }, options)
}
#[derive(Debug, Clone, Copy)]
enum Method<'body> {
Get,
Post {
content_type: &'body str,
body: &'body [u8],
},
}
const CREDENTIAL_HEADERS: [&str; 3] = ["Authorization", "Cookie", "Proxy-Authorization"];
#[derive(Debug, PartialEq, Eq)]
struct Origin {
scheme: String,
host: String,
port: Option<u16>,
}
fn origin_of(target: &str) -> Option<Origin> {
let target = UriAbsoluteStr::new(target).ok()?;
let authority = target.authority_components()?;
let scheme = target.scheme_str().to_ascii_lowercase();
let port = match authority.port() {
Some(port) if !port.is_empty() => Some(port.parse::<u16>().ok()?),
_ => match scheme.as_str() {
"http" => Some(80),
"https" => Some(443),
_ => None,
},
};
Some(Origin {
scheme,
host: authority.host().to_ascii_lowercase(),
port,
})
}
fn resolve_location(base: &str, location: &str) -> Result<String, Error> {
let refused = |detail: &str| {
failure(
FailureStage::Redirect,
FailureKind::Transport,
Some(detail.to_owned()),
)
};
let base = match UriAbsoluteStr::new(base) {
Ok(base) => base,
Err(_not_absolute) => {
let refusal = Err(refused("redirect base is not absolute"));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "resolve_location: returning an error to the caller");
return refusal;
}
};
let reference = match UriReferenceStr::new(location) {
Ok(reference) => reference,
Err(_not_a_reference) => {
let refusal = Err(refused("Location is not a URI reference"));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "resolve_location: returning an error to the caller");
return refusal;
}
};
let resolved = reference.resolve_against(base);
if resolved.ensure_rfc3986_normalizable().is_err() {
let refusal = Err(refused(
"Location does not resolve to one unambiguous target",
));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "resolve_location: returning an error to the caller");
return refusal;
}
let resolved = match resolved.try_to_dedicated_string() {
Ok(resolved) => resolved,
Err(_not_representable) => {
let refusal = Err(failure(FailureStage::Redirect, FailureKind::Resource, None));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "resolve_location: returning an error to the caller");
return refusal;
}
};
let target = resolved
.as_str()
.split_once('#')
.map_or(resolved.as_str(), |(target, _fragment)| target);
if validate_url(target).is_err() {
let refusal = Err(refused("Location leaves absolute http(s)"));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "resolve_location: returning an error to the caller");
return refusal;
}
Ok(target.to_owned())
}
fn redirect_location(response: &ureq::http::Response<ureq::Body>) -> Result<Option<&str>, Error> {
let status = response.status();
if !status.is_redirection() || status == ureq::http::StatusCode::NOT_MODIFIED {
return Ok(None);
}
let Some(location) = response
.headers()
.get_all(ureq::http::header::LOCATION)
.iter()
.next_back()
else {
return Ok(None);
};
match location.to_str() {
Ok(text) => Ok(Some(text)),
Err(_not_visible_ascii) => Err(failure_unprintable_location(location.as_bytes())),
}
}
fn next_method<'body>(
status: ureq::http::StatusCode,
method: Method<'body>,
) -> Result<Method<'body>, Error> {
let keeps_method = matches!(status.as_u16(), 307 | 308);
match method {
Method::Post { .. } if keeps_method => Err(failure(
FailureStage::Redirect,
FailureKind::Transport,
Some("a 307/308 redirect would resend the request body; not followed".to_owned()),
)),
Method::Post { .. } | Method::Get => Ok(Method::Get),
}
}
fn send_hop<'headers>(
agent: &ureq::Agent,
target: &str,
method: Method<'_>,
headers: impl Iterator<Item = &'headers (String, String)>,
remaining: Option<Duration>,
) -> Result<ureq::http::Response<ureq::Body>, Error> {
let result = match method {
Method::Get => headers
.fold(
agent.get(target).config().timeout_global(remaining).build(),
|call, header| call.header(header.0.as_str(), header.1.as_str()),
)
.call(),
Method::Post { content_type, body } => headers
.fold(
agent
.post(target)
.config()
.timeout_global(remaining)
.build()
.header("Content-Type", content_type),
|call, header| call.header(header.0.as_str(), header.1.as_str()),
)
.send(body),
};
result.map_err(|error| map_error(error, is_https(target)))
}
fn exchange(url: &str, method: Method<'_>, options: &Options) -> Result<Response, Error> {
validate_url(url)?;
let limit = redirect_limit(options)?;
let agent = agent(options);
let origin = origin_of(url);
let mut target = url.to_owned();
let mut method = method;
let mut chain = vec![sanitized_target(url)];
let mut hops: u32 = 0;
let deadline_at = options
.deadline
.and_then(|deadline| std::time::Instant::now().checked_add(deadline));
loop {
let remaining = match deadline_at {
Some(at) => match at.checked_duration_since(std::time::Instant::now()) {
Some(left) if !left.is_zero() => Some(left),
_ => {
{
let refusal =
Err(failure(FailureStage::Deadline, FailureKind::Timeout, None));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "exchange: returning an error to the caller");
return refusal;
};
}
},
None => None,
};
let first = hops == 0;
let same_origin = origin.is_some() && origin_of(&target) == origin;
let headers = options.headers.iter().filter(|header| {
first
|| (same_origin
&& !CREDENTIAL_HEADERS
.iter()
.any(|credential| header.0.eq_ignore_ascii_case(credential)))
});
let response = send_hop(&agent, &target, method, headers, remaining)?;
if limit == 0 {
return response_of(response, options, chain);
}
match next_hop(&target, &response, limit, hops, method)? {
Hop::Final => return response_of(response, options, chain),
Hop::Follow {
next,
method: next_method,
hops: climbed,
} => {
method = next_method;
hops = climbed;
chain.push(sanitized_target(&next));
target = next;
}
}
}
}
enum Hop<'body> {
Final,
Follow {
next: String,
method: Method<'body>,
hops: u32,
},
}
fn next_hop<'body>(
current: &str,
response: &ureq::http::Response<ureq::Body>,
limit: u32,
hops: u32,
method: Method<'body>,
) -> Result<Hop<'body>, Error> {
let Some(location) = redirect_location(response)? else {
return Ok(Hop::Final);
};
if hops >= limit {
let refusal = Err(failure(
FailureStage::Redirect,
FailureKind::RedirectLimit,
None,
));
#[cfg(feature = "trace")]
crate::trace::debug!(error = ?refusal.as_ref().err(), "next_hop: returning an error to the caller");
return refusal;
}
let next = resolve_location(current, location)?;
let method = next_method(response.status(), method)?;
Ok(Hop::Follow {
next,
method,
hops: hops.saturating_add(1),
})
}
#[cfg(test)]
#[expect(
clippy::disallowed_methods,
reason = "loopback test servers need a real thread, and holding one open past a read timeout needs a real sleep"
)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::thread;
const ECHO: &str = "echo-body-123";
fn read_request_head(stream: &mut TcpStream) -> std::io::Result<Vec<u8>> {
let mut request = [0_u8; 1024];
let mut head = Vec::new();
while !head.windows(4).any(|window| window == b"\r\n\r\n") {
let read = stream.read(&mut request)?;
if read == 0 {
break;
}
head.extend_from_slice(&request[..read]);
}
Ok(head)
}
fn serve_held_response(
response: Vec<u8>,
) -> std::io::Result<(
u16,
std::sync::mpsc::SyncSender<()>,
thread::JoinHandle<std::io::Result<()>>,
)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let (release, wait_release) = std::sync::mpsc::sync_channel(1);
let handle = thread::spawn(move || -> std::io::Result<()> {
let (mut stream, _) = listener.accept()?;
let _ = read_request_head(&mut stream)?;
stream.write_all(&response)?;
let _released = wait_release.recv();
Ok(())
});
Ok((port, release, handle))
}
fn serve_held_socket() -> std::io::Result<(
u16,
std::sync::mpsc::SyncSender<()>,
thread::JoinHandle<std::io::Result<()>>,
)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let (release, wait_release) = std::sync::mpsc::sync_channel(1);
let handle = thread::spawn(move || -> std::io::Result<()> {
let (_stream, _) = listener.accept()?;
let _released = wait_release.recv();
Ok(())
});
Ok((port, release, handle))
}
fn serve(
replies: Vec<(&'static str, String)>,
) -> std::io::Result<(u16, thread::JoinHandle<std::io::Result<()>>)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let handle = thread::spawn(move || -> std::io::Result<()> {
for (status, body) in replies {
let (mut stream, _) = listener.accept()?;
serve_one(&mut stream, status, body)?;
}
Ok(())
});
Ok((port, handle))
}
#[cfg(test)]
fn serve_one(
stream: &mut std::net::TcpStream,
status: &'static str,
body: String,
) -> std::io::Result<()> {
let mut request = vec![0u8; 4096];
let mut head = Vec::new();
loop {
let n = stream.read(&mut request)?;
head.extend_from_slice(&request[..n]);
if head.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header_end = head
.windows(4)
.position(|w| w == b"\r\n\r\n")
.map(|position| position.saturating_add(4))
.unwrap_or(head.len());
let text = String::from_utf8_lossy(&head[..header_end]);
let content_length = text
.lines()
.filter_map(|line| line.split_once(':'))
.find(|entry| entry.0.eq_ignore_ascii_case("content-length"))
.and_then(|(_, value)| value.trim().parse::<usize>().ok())
.unwrap_or(0);
let mut received = head.len().saturating_sub(header_end);
while received < content_length {
let n = stream.read(&mut request)?;
head.extend_from_slice(&request[..n]);
received = received.saturating_add(n);
}
let full = String::from_utf8_lossy(&head);
let echoed = full.contains(ECHO);
let payload = if echoed { ECHO.to_owned() } else { body };
let reply = format!(
"HTTP/1.1 {status}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{payload}",
payload.len()
);
stream.write_all(reply.as_bytes())
}
fn join_server(
server: thread::JoinHandle<std::io::Result<()>>,
) -> Result<(), Box<dyn std::error::Error>> {
let served = server
.join()
.map_err(|_| "the canned server thread panicked before replying")?;
served?;
Ok(())
}
fn serve_raw(
reply: Vec<u8>,
) -> std::io::Result<(u16, thread::JoinHandle<std::io::Result<()>>)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let handle = thread::spawn(move || -> std::io::Result<()> {
let (mut stream, _) = listener.accept()?;
let mut request = vec![0u8; 4096];
let mut head = Vec::new();
loop {
let n = stream.read(&mut request)?;
head.extend_from_slice(&request[..n]);
if head.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
stream.write_all(&reply)
});
Ok((port, handle))
}
#[test]
fn default_option_wrappers_reach_the_same_path() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![
("200 OK", "hello".to_owned()),
("200 OK", "ok".to_owned()),
])?;
let url = format!("http://127.0.0.1:{port}/");
let got = get(&url)?;
assert_eq!(got.status, 200);
assert_eq!(got.body, b"hello");
let posted = post(&url, "text/plain", ECHO.as_bytes())?;
assert_eq!(posted.status, 200);
assert_eq!(posted.text()?, ECHO);
join_server(server)?;
Ok(())
}
fn quiet() -> Options {
Options {
timeout: Duration::from_secs(5),
deadline: None,
user_agent: "lgwks-std-test".into(),
headers: Vec::new(),
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
body_policy: BodyPolicy::Whole,
redirect_policy: RedirectPolicy::default(),
}
}
const SMALL_CEILING: usize = 64;
fn ceiling(max_body_bytes: usize) -> Options {
quiet().max_body_bytes(max_body_bytes)
}
fn previewing(max_body_bytes: usize) -> Options {
ceiling(max_body_bytes).body_policy(BodyPolicy::Preview)
}
fn filler(bytes: usize) -> String {
"x".repeat(bytes)
}
#[test]
fn gets_status_headers_and_body() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![("200 OK", "hello".to_owned())])?;
let response = get_with(&format!("http://127.0.0.1:{port}/"), &quiet())?;
assert_eq!(response.status, 200);
assert_eq!(response.body, b"hello");
assert_eq!(response.text()?, "hello");
assert!(
response
.headers
.iter()
.any(|header| header.0.eq_ignore_ascii_case("content-length"))
);
join_server(server)?;
Ok(())
}
#[test]
fn error_statuses_are_responses_not_errors() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![("404 Not Found", "missing".to_owned())])?;
let response = get_with(&format!("http://127.0.0.1:{port}/"), &quiet())?;
assert_eq!(response.status, 404);
assert_eq!(response.text()?, "missing");
join_server(server)?;
Ok(())
}
#[test]
fn legal_header_bytes_and_repeated_values_are_preserved()
-> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_raw(
b"HTTP/1.1 200 OK\r\nX-Value: \x80A\r\nX-Value: \x81A\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.to_vec(),
)?;
let response = get_with(&format!("http://127.0.0.1:{port}/"), &quiet())?;
let values: Vec<_> = response
.header_values()
.filter(|header| header.0.eq_ignore_ascii_case("x-value"))
.map(|header| header.1.to_vec())
.collect();
assert_eq!(values.len(), 2, "both repeated values are retained");
assert!(
values.contains(&vec![0x80, b'A']),
"first raw value is exact"
);
assert!(
values.contains(&vec![0x81, b'A']),
"second raw value is exact"
);
let display_values: Vec<_> = response
.headers()
.iter()
.filter(|header| header.0.eq_ignore_ascii_case("x-value"))
.map(|header| header.1.clone())
.collect();
assert_eq!(
display_values.len(),
2,
"the compatibility view keeps multiplicity"
);
assert_eq!(
display_values[0], display_values[1],
"lossy display is explicitly lossy"
);
join_server(server)?;
Ok(())
}
#[test]
fn no_follow_returns_the_redirect_without_contacting_its_target()
-> Result<(), Box<dyn std::error::Error>> {
let reply = b"HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:9/next?token=secret\r\nContent-Length: 0\r\nConnection: close\r\n\r\n";
let (port, server) = serve_raw(reply.to_vec())?;
let response = get_with(
&format!("http://127.0.0.1:{port}/start?credential=hidden"),
&quiet().redirect_policy(RedirectPolicy::NoFollow),
)?;
assert_eq!(
response.status, 302,
"the redirect remains an HTTP response"
);
assert_eq!(
response.final_target(),
format!("http://127.0.0.1:{port}/start")
);
assert_eq!(response.redirect_chain().len(), 1);
assert!(!response.final_target().contains("hidden"));
join_server(server)?;
Ok(())
}
#[test]
fn followed_redirect_records_target_and_strips_sensitive_headers()
-> Result<(), Box<dyn std::error::Error>> {
let destination = TcpListener::bind("127.0.0.1:0")?;
let destination_port = destination.local_addr()?.port();
let destination_server = thread::spawn(move || -> std::io::Result<String> {
let (mut stream, _) = destination.accept()?;
let request = read_request_head(&mut stream)?;
stream.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok",
)?;
Ok(String::from_utf8_lossy(&request).into_owned())
});
let origin = TcpListener::bind("127.0.0.1:0")?;
let origin_port = origin.local_addr()?.port();
let origin_server = thread::spawn(move || -> std::io::Result<()> {
let (mut stream, _) = origin.accept()?;
let _ = read_request_head(&mut stream)?;
let reply = format!(
"HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:{destination_port}/final?token=hidden\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
stream.write_all(reply.as_bytes())
});
let options = quiet()
.header("Authorization", "Bearer secret")
.redirect_policy(RedirectPolicy::FollowAtMost(1));
let response = get_with(
&format!("http://127.0.0.1:{origin_port}/start?token=hidden"),
&options,
)?;
assert_eq!(response.status, 200);
assert_eq!(
response.final_target(),
format!("http://127.0.0.1:{destination_port}/final")
);
assert_eq!(response.redirect_chain().len(), 2);
assert!(
response
.redirect_chain()
.iter()
.all(|target| !target.contains("hidden") && !target.contains("secret")),
"redirect provenance strips query secrets"
);
let received = destination_server
.join()
.map_err(|_| "redirect destination thread panicked")??;
assert!(
!received.to_ascii_lowercase().contains("authorization:"),
"sensitive authorization is not forwarded"
);
origin_server
.join()
.map_err(|_| "redirect origin thread panicked")??;
Ok(())
}
#[test]
fn redirect_loop_refuses_at_the_configured_limit() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let server = thread::spawn(move || -> std::io::Result<usize> {
let mut count = 0;
for _ in 0..=2 {
let (mut stream, _) = listener.accept()?;
let _ = read_request_head(&mut stream)?;
let reply = format!(
"HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:{port}/loop\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
stream.write_all(reply.as_bytes())?;
count += 1;
}
Ok(count)
});
let result = get_with(
&format!("http://127.0.0.1:{port}/loop"),
&quiet().redirect_policy(RedirectPolicy::FollowAtMost(2)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Redirect,
kind: FailureKind::RedirectLimit,
..
})
));
assert_eq!(
server
.join()
.map_err(|_| "redirect loop server panicked")??,
3,
"the receiver observed the initial request plus two followed hops"
);
Ok(())
}
#[test]
fn a_deadline_bounds_the_whole_redirect_chain() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let server = thread::spawn(move || -> std::io::Result<usize> {
let mut served = 0;
listener.set_nonblocking(true)?;
let quiet_since = std::time::Instant::now();
while quiet_since.elapsed() < Duration::from_secs(3) && served < 10 {
let mut stream = match listener.accept() {
Ok((stream, _)) => stream,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(5));
continue;
}
Err(error) => return Err(error),
};
stream.set_nonblocking(false)?;
let _ = read_request_head(&mut stream)?;
thread::sleep(Duration::from_millis(200));
let reply = format!(
"HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:{port}/hop\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
if stream.write_all(reply.as_bytes()).is_err() {
served += 1;
break;
}
served += 1;
}
Ok(served)
});
let started = std::time::Instant::now();
let result = get_with(
&format!("http://127.0.0.1:{port}/hop"),
&quiet()
.timeout(Duration::from_secs(2))
.deadline(Duration::from_millis(500))
.redirect_policy(RedirectPolicy::FollowAtMost(10)),
);
let elapsed = started.elapsed();
assert!(
matches!(
result,
Err(Error::Failure {
stage: FailureStage::Deadline,
kind: FailureKind::Timeout,
..
})
),
"the chain is refused at its deadline, got {result:?}"
);
assert!(
elapsed < Duration::from_millis(1_500),
"the deadline bounds the chain, not ten hops of 200 ms: {elapsed:?}"
);
let served = server.join().map_err(|_| "redirect server panicked")??;
assert!(
(2..10).contains(&served),
"the chain made progress, then stopped at the deadline: {served} hops"
);
Ok(())
}
#[test]
fn a_deadline_bounds_a_trickled_body() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let server = thread::spawn(move || -> std::io::Result<()> {
let (mut stream, _) = listener.accept()?;
let _ = read_request_head(&mut stream)?;
stream.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 64\r\n\r\n")?;
for _ in 0..64 {
thread::sleep(Duration::from_millis(50));
if stream.write_all(b"x").is_err() {
break;
}
}
Ok(())
});
let started = std::time::Instant::now();
let result = get_with(
&format!("http://127.0.0.1:{port}/"),
&quiet()
.timeout(Duration::from_secs(2))
.deadline(Duration::from_millis(400)),
);
let elapsed = started.elapsed();
assert!(
matches!(
result,
Err(Error::Failure {
stage: FailureStage::Deadline,
kind: FailureKind::Timeout,
..
})
),
"a trickled body is stopped at the call deadline, got {result:?}"
);
assert!(
elapsed < Duration::from_millis(1_500),
"the body was not read to its 3.2 s end: {elapsed:?}"
);
server.join().map_err(|_| "trickle server panicked")??;
Ok(())
}
#[test]
fn body_capacity_grows_geometrically_within_the_ceiling()
-> Result<(), Box<dyn std::error::Error>> {
let body_len = 1 << 20;
let ceiling = body_len + 3;
let mut body = Vec::new();
let mut growths = 0;
let mut last_capacity = body.capacity();
while body.len() < body_len {
let incoming = READ_CHUNK_BYTES.min(body_len - body.len());
reserve_for(&mut body, incoming, ceiling)?;
body.extend(std::iter::repeat_n(b'x', incoming));
if body.capacity() != last_capacity {
growths += 1;
last_capacity = body.capacity();
}
assert!(
body.capacity() <= ceiling,
"capacity stays within the ceiling"
);
}
assert!(growths <= 9, "{growths} reallocations for 128 reads");
let mut near_limit = Vec::with_capacity(10);
near_limit.extend([0_u8; 10]);
reserve_for(&mut near_limit, 5, 15)?;
assert!(near_limit.capacity() >= 15, "room for the incoming read");
assert!(
near_limit.capacity() < 20,
"the reservation was clamped to the ceiling: {}",
near_limit.capacity()
);
Ok(())
}
#[test]
fn oversized_redirect_limit_is_refused_before_dialing() {
assert!(matches!(
get_with(
"http://127.0.0.1:9/",
&quiet().redirect_policy(RedirectPolicy::FollowAtMost(
MAX_REDIRECT_HOPS.saturating_add(1)
)),
),
Err(Error::RedirectLimitTooLarge { .. })
));
}
const OK_REPLY: &str = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok";
fn redirect_reply(status: &str, location: &str) -> String {
format!(
"HTTP/1.1 {status}\r\nLocation: {location}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
}
fn serve_recording(
replies: Vec<String>,
) -> std::io::Result<(u16, thread::JoinHandle<std::io::Result<Vec<String>>>)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let handle = thread::spawn(move || -> std::io::Result<Vec<String>> {
let mut heads = Vec::new();
for reply in replies {
let (mut stream, _) = listener.accept()?;
let head = read_request_head(&mut stream)?;
heads.push(String::from_utf8_lossy(&head).to_ascii_lowercase());
stream.write_all(reply.as_bytes())?;
}
Ok(heads)
});
Ok((port, handle))
}
#[test]
fn a_cross_origin_redirect_carries_no_caller_header() -> Result<(), Box<dyn std::error::Error>>
{
let (destination_port, destination) = serve_recording(vec![OK_REPLY.to_owned()])?;
let (origin_port, origin) = serve_recording(vec![redirect_reply(
"302 Found",
&format!("http://127.0.0.1:{destination_port}/landed"),
)])?;
let options = quiet()
.header("X-Api-Key", "key-secret")
.header("Proxy-Authorization", "Basic proxy-secret")
.header("Authorization", "Bearer bearer-secret");
let response = get_with(&format!("http://127.0.0.1:{origin_port}/start"), &options)?;
assert_eq!(response.status, 200);
let sent = origin
.join()
.map_err(|_| "origin server panicked")??
.concat();
assert!(
sent.contains("x-api-key: key-secret"),
"the first hop carries the caller's headers:\n{sent}"
);
let forwarded = destination
.join()
.map_err(|_| "destination server panicked")??
.concat();
assert!(
!forwarded.contains("secret"),
"no caller header reaches another origin:\n{forwarded}"
);
Ok(())
}
#[test]
fn a_same_origin_redirect_keeps_ordinary_headers_but_not_credentials()
-> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_recording(vec![
redirect_reply("302 Found", "next?page=2#section"),
OK_REPLY.to_owned(),
])?;
let options = quiet()
.header("X-Api-Key", "key-value")
.header("Authorization", "Bearer bearer-secret");
let response = get_with(&format!("http://127.0.0.1:{port}/dir/start"), &options)?;
assert_eq!(
response.final_target(),
format!("http://127.0.0.1:{port}/dir/next")
);
let heads = server.join().map_err(|_| "server panicked")??;
let (Some(first), Some(second), 2) = (heads.first(), heads.get(1), heads.len()) else {
return Err(format!("expected two requests, saw {}", heads.len()).into());
};
assert!(first.contains("authorization: bearer bearer-secret"));
assert!(
second.starts_with("get /dir/next?page=2 "),
"the relative Location resolves against the target, without its fragment:\n{second}"
);
assert!(
second.contains("x-api-key: key-value"),
"same origin keeps ordinary headers:\n{second}"
);
assert!(
!second.contains("authorization:"),
"credentials are not re-sent on a server's say-so:\n{second}"
);
Ok(())
}
#[test]
fn a_see_other_redirect_turns_a_post_into_a_bodiless_get()
-> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_recording(vec![
redirect_reply("303 See Other", "/after"),
OK_REPLY.to_owned(),
])?;
let response = post_with(
&format!("http://127.0.0.1:{port}/submit"),
"application/json",
b"",
&quiet(),
)?;
assert_eq!(response.status, 200);
let heads = server.join().map_err(|_| "server panicked")??;
let (Some(first), Some(second), 2) = (heads.first(), heads.get(1), heads.len()) else {
return Err(format!("expected two requests, saw {}", heads.len()).into());
};
assert!(first.starts_with("post /submit "), "{first}");
assert!(second.starts_with("get /after "), "{second}");
assert!(
!second.contains("content-type:"),
"the GET carries no body, so no content type:\n{second}"
);
Ok(())
}
#[test]
fn a_method_keeping_redirect_of_a_post_is_refused() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) =
serve_recording(vec![redirect_reply("307 Temporary Redirect", "/again")])?;
let result = post_with(
&format!("http://127.0.0.1:{port}/submit"),
"application/json",
b"",
&quiet(),
);
assert!(
matches!(
result,
Err(Error::Failure {
stage: FailureStage::Redirect,
kind: FailureKind::Transport,
..
})
),
"{result:?}"
);
let heads = server.join().map_err(|_| "server panicked")??;
assert_eq!(heads.len(), 1, "the body was not replayed");
Ok(())
}
#[test]
fn origins_compare_scheme_host_and_effective_port() {
assert!(origin_of("HTTP://Example.COM/a").is_some());
assert_eq!(
origin_of("HTTP://Example.COM/a"),
origin_of("http://example.com:80/b")
);
assert_eq!(
origin_of("https://example.com/"),
origin_of("https://example.com:443/")
);
assert_ne!(
origin_of("https://example.com/"),
origin_of("http://example.com/"),
"a downgrade is another origin"
);
assert_ne!(
origin_of("http://example.com:8080/"),
origin_of("http://example.com/")
);
assert_eq!(
origin_of("http://[::1]:8080/"),
origin_of("http://[::1]:8080/x")
);
assert!(
origin_of("http://example.com:99999/").is_none(),
"an unreadable port is never the same origin"
);
}
#[test]
fn a_location_resolves_to_an_absolute_http_target_or_is_refused() -> Result<(), Error> {
assert_eq!(
resolve_location("http://a.test/dir/page", "next?x=1#top")?,
"http://a.test/dir/next?x=1"
);
assert_eq!(
resolve_location("http://a.test/dir/page", "//b.test/p")?,
"http://b.test/p"
);
for hostile in [
"ftp://a.test/token-secret",
"javascript:token-secret",
"http://",
"not a uri token-secret",
] {
let refusal = resolve_location("http://a.test/", hostile);
assert!(
matches!(
refusal,
Err(Error::Failure {
stage: FailureStage::Redirect,
..
})
),
"{hostile} must be refused"
);
assert!(
!format!("{refusal:?}").contains("token-secret"),
"the Location is not echoed"
);
}
Ok(())
}
#[test]
fn posts_body_with_content_type() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![("200 OK", String::new())])?;
let response = post_with(
&format!("http://127.0.0.1:{port}/"),
"text/plain",
ECHO.as_bytes(),
&quiet(),
)?;
assert_eq!(response.status, 200);
assert_eq!(response.text()?, ECHO);
join_server(server)?;
Ok(())
}
#[test]
fn rejects_non_http_urls_before_dialing() {
assert!(matches!(
get_with("not a url", &quiet()),
Err(Error::InvalidUrl)
));
assert!(matches!(
get_with("/relative/path", &quiet()),
Err(Error::InvalidUrl)
));
assert!(matches!(
get_with("ftp://127.0.0.1/file", &quiet()),
Err(Error::InvalidUrl)
));
assert!(matches!(
get_with("http:user:SECRET@host", &quiet()),
Err(Error::InvalidUrl)
));
assert!(matches!(
get_with("https://", &quiet()),
Err(Error::InvalidUrl)
));
}
#[test]
fn refused_connection_is_transport_not_timeout() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
let Err(error) = get_with(&format!("http://127.0.0.1:{port}/"), &quiet()) else {
return Err("a refused connection must not yield a response".into());
};
assert!(matches!(
error,
Error::Failure {
stage: FailureStage::Connect,
kind: FailureKind::Transport,
..
}
));
Ok(())
}
#[test]
fn custom_headers_reach_the_server() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let handle = thread::spawn(move || -> std::io::Result<String> {
let (mut stream, _) = listener.accept()?;
let mut request = vec![0u8; 4096];
let mut head = Vec::new();
loop {
let n = stream.read(&mut request)?;
head.extend_from_slice(&request[..n]);
if head.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let text = String::from_utf8_lossy(&head).into_owned();
let reply = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok";
stream.write_all(reply.as_bytes())?;
Ok(text)
});
let options = quiet()
.header("Authorization", "Bearer test-token")
.idempotency_key("old-operation")
.idempotency_key("current-operation");
let response = post_with(
&format!("http://127.0.0.1:{port}/"),
"text/plain",
b"hi",
&options,
)?;
assert_eq!(response.status, 200);
let seen = handle
.join()
.map_err(|_| "the recording server thread panicked before replying")??;
assert!(
seen.to_ascii_lowercase()
.contains("authorization: bearer test-token"),
"server never saw the Authorization header:\n{seen}"
);
assert_eq!(
seen.to_ascii_lowercase()
.matches("idempotency-key:")
.count(),
1,
"the receiver sees exactly one idempotency key"
);
assert!(
seen.to_ascii_lowercase()
.contains("idempotency-key: current-operation")
);
Ok(())
}
#[test]
fn silent_server_hits_timeout() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let (release, released) = std::sync::mpsc::sync_channel(1);
let handle = thread::spawn(move || -> std::io::Result<()> {
let (_stream, _) = listener.accept()?;
let _released = released.recv();
Ok(())
});
let options = quiet().timeout(Duration::from_millis(200));
let Err(error) = get_with(&format!("http://127.0.0.1:{port}/"), &options) else {
return Err("a silent server must hit the read timeout".into());
};
assert!(
matches!(
error,
Error::Failure {
stage: FailureStage::Headers,
kind: FailureKind::Timeout,
..
} | Error::Failure {
kind: FailureKind::Interrupted,
..
}
),
"a silent server yields a header timeout or an interruption, not {error:?}"
);
release
.send(())
.map_err(|_| "timeout fixture receiver ended")?;
handle
.join()
.map_err(|_| "timeout fixture thread panicked")??;
Ok(())
}
#[test]
fn tls_handshake_timeout_preserves_connect_stage() -> Result<(), Box<dyn std::error::Error>> {
let (port, release, server) = serve_held_socket()?;
let result = get_with(
&format!("https://127.0.0.1:{port}/"),
&quiet().timeout(Duration::from_millis(150)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Connect,
kind: FailureKind::Timeout,
..
})
));
release.send(()).map_err(|_| "TLS fixture receiver ended")?;
server.join().map_err(|_| "TLS fixture thread panicked")??;
Ok(())
}
#[test]
fn malformed_tls_handshake_preserves_tls_stage() -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let server = thread::spawn(move || -> std::io::Result<()> {
let (mut stream, _) = listener.accept()?;
let mut hello = [0_u8; 1024];
let _hello_len = stream.read(&mut hello)?;
stream.write_all(b"not-a-tls-record")?;
let _hung_up = std::io::copy(&mut stream, &mut std::io::sink());
Ok(())
});
let result = get_with(
&format!("https://127.0.0.1:{port}/"),
&quiet().timeout(Duration::from_secs(2)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Tls,
kind: FailureKind::Transport,
..
})
));
server
.join()
.map_err(|_| "malformed TLS fixture panicked")??;
Ok(())
}
#[test]
fn request_body_timeout_preserves_send_stage() -> Result<(), Box<dyn std::error::Error>> {
let (port, release, server) = serve_held_socket()?;
let body = vec![b'x'; 32 * 1024 * 1024];
let result = post_with(
&format!("http://127.0.0.1:{port}/"),
"application/octet-stream",
&body,
&quiet().timeout(Duration::from_millis(150)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Send,
kind: FailureKind::Timeout,
..
})
));
release
.send(())
.map_err(|_| "send fixture receiver ended")?;
server
.join()
.map_err(|_| "send fixture thread panicked")??;
Ok(())
}
#[test]
fn body_timeout_preserves_stage_and_class() -> Result<(), Box<dyn std::error::Error>> {
let (port, release, server) = serve_held_response(
b"HTTP/1.1 200 OK\r\nContent-Length: 8\r\nConnection: keep-alive\r\n\r\nx".to_vec(),
)?;
let result = get_with(
&format!("http://127.0.0.1:{port}/"),
&quiet().timeout(Duration::from_millis(150)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Body,
kind: FailureKind::Timeout,
..
})
));
release
.send(())
.map_err(|_| "body fixture receiver ended")?;
server
.join()
.map_err(|_| "body fixture thread panicked")??;
Ok(())
}
#[test]
fn eof_probe_timeout_preserves_stage_and_class() -> Result<(), Box<dyn std::error::Error>> {
let limit = SMALL_CEILING;
let mut response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: keep-alive\r\n\r\n",
limit.saturating_add(1)
)
.into_bytes();
response.extend(std::iter::repeat_n(b'x', limit));
let (port, release, server) = serve_held_response(response)?;
let result = get_with(
&format!("http://127.0.0.1:{port}/"),
&previewing(limit).timeout(Duration::from_millis(150)),
);
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::EofProbe,
kind: FailureKind::Timeout,
..
})
));
release.send(()).map_err(|_| "EOF fixture receiver ended")?;
server.join().map_err(|_| "EOF fixture thread panicked")??;
Ok(())
}
#[test]
fn truncated_transport_is_not_reported_as_complete_eof()
-> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Length: 8\r\nConnection: close\r\n\r\nabc".to_vec(),
)?;
let result = get_with(&format!("http://127.0.0.1:{port}/"), &quiet());
assert!(matches!(
result,
Err(Error::Failure {
stage: FailureStage::Body,
kind: FailureKind::Transport,
..
})
));
join_server(server)?;
Ok(())
}
#[test]
fn invalid_utf8_is_a_payload_decoding_failure() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Length: 1\r\nConnection: close\r\n\r\n\xff".to_vec(),
)?;
let response = get_with(&format!("http://127.0.0.1:{port}/"), &quiet())?;
assert_eq!(response.status, 200, "the HTTP exchange completed");
assert!(matches!(
response.text(),
Err(Error::Failure {
stage: FailureStage::TextDecode,
kind: FailureKind::InvalidUtf8,
..
})
));
join_server(server)?;
Ok(())
}
#[test]
fn non_power_of_two_ceiling_bounds_retained_capacity() -> Result<(), Box<dyn std::error::Error>>
{
let limit = 73;
let mut reader = std::io::Cursor::new(filler(limit).into_bytes());
let (body, truncation) = read_bounded(&mut reader, &ceiling(limit))?;
assert_eq!(body.len(), limit, "the exact body is retained");
assert!(
body.capacity() <= limit,
"body capacity stays within the declared ceiling"
);
assert_eq!(truncation, Truncation::Complete);
Ok(())
}
#[test]
fn the_read_window_is_a_fixed_chunk_and_capacity_is_clamped_to_the_ceiling()
-> Result<(), Box<dyn std::error::Error>> {
struct WindowRecorder {
remaining: usize,
largest_window: usize,
reads: usize,
}
impl Read for WindowRecorder {
fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
self.largest_window = self.largest_window.max(buffer.len());
self.reads = self.reads.saturating_add(1);
if buffer.is_empty() || self.remaining == 0 {
return Ok(0);
}
buffer[0] = b'x';
self.remaining = self.remaining.saturating_sub(1);
Ok(1)
}
}
for limit in [73_usize, 1000, 3003, 10_000] {
let mut reader = WindowRecorder {
remaining: limit,
largest_window: 0,
reads: 0,
};
let (body, truncation) = read_bounded(&mut reader, &ceiling(limit))?;
assert_eq!(
body.len(),
limit,
"the exact body is retained at a {limit}-byte ceiling"
);
assert_eq!(
body.capacity(),
limit,
"capacity is clamped to the non-power-of-two ceiling at {limit}: {}",
body.capacity()
);
assert_eq!(
truncation,
Truncation::Complete,
"an exactly-at-ceiling body ends on its own at {limit}"
);
assert!(
reader.largest_window <= READ_CHUNK_BYTES,
"the reader window is one fixed chunk at {limit}, not the body: {}",
reader.largest_window
);
if limit > READ_CHUNK_BYTES {
assert_eq!(
reader.largest_window, READ_CHUNK_BYTES,
"the first window of a body larger than one chunk is exactly one chunk"
);
}
assert_eq!(
reader.reads,
limit.saturating_add(1),
"one read per byte plus the single EOF probe at {limit}"
);
}
Ok(())
}
#[test]
fn an_interruption_is_classified_the_same_way_at_every_stage() {
let interrupted = || std::io::Error::from(std::io::ErrorKind::Interrupted);
assert!(
matches!(
map_error(ureq::Error::Io(interrupted()), false),
Error::Failure {
kind: FailureKind::Interrupted,
..
}
),
"EINTR on the request path must not be reported as a transport failure"
);
for stage in [FailureStage::Body, FailureStage::EofProbe] {
assert!(
matches!(
map_read_error(interrupted(), stage),
Error::Failure {
kind: FailureKind::Interrupted,
stage: observed,
..
} if observed == stage
),
"EINTR while reading at {stage:?} keeps its stage and its class"
);
}
assert!(matches!(
map_error(
ureq::Error::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)),
false,
),
Error::Failure {
kind: FailureKind::Transport,
..
}
));
}
#[test]
fn the_https_scheme_is_recognised_in_any_case() {
assert!(is_https("https://example.com/"));
assert!(is_https("HTTPS://example.com/"));
assert!(is_https("HtTpS://example.com/"));
assert!(!is_https("http://example.com/"));
assert!(!is_https("https:"));
}
#[test]
fn a_body_under_the_ceiling_is_whole() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![("200 OK", "hello".to_owned())])?;
let response = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
)?;
assert_eq!(response.body, b"hello");
assert_eq!(
response.truncation,
Truncation::Complete,
"a body below the ceiling ended on its own"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_body_exactly_at_the_ceiling_is_whole() -> Result<(), Box<dyn std::error::Error>> {
let payload = filler(SMALL_CEILING);
let (port, server) = serve(vec![("200 OK", payload)])?;
let response = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
)?;
assert_eq!(
response.body.len(),
SMALL_CEILING,
"the whole body is kept at the ceiling"
);
assert_eq!(
response.truncation,
Truncation::Complete,
"a body that ends at the ceiling has not overflowed it"
);
join_server(server)?;
Ok(())
}
#[test]
fn zero_ceiling_accepts_an_empty_body() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n".to_vec(),
)?;
let response = get_with(&format!("http://127.0.0.1:{port}/"), &ceiling(0))?;
assert!(
response.body.is_empty(),
"zero ceiling retains no body bytes"
);
assert_eq!(response.truncation, Truncation::Complete);
join_server(server)?;
Ok(())
}
#[test]
fn zero_ceiling_refuses_a_nonempty_whole_body() -> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Length: 1\r\nConnection: close\r\n\r\nx".to_vec(),
)?;
let result = get_with(&format!("http://127.0.0.1:{port}/"), &ceiling(0));
assert!(matches!(result, Err(Error::BodyTooLarge { limit: 0 })));
join_server(server)?;
Ok(())
}
#[test]
fn a_body_one_byte_past_the_ceiling_is_refused() -> Result<(), Box<dyn std::error::Error>> {
let payload = filler(SMALL_CEILING.saturating_add(1));
let (port, server) = serve(vec![("200 OK", payload)])?;
let Err(error) = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
) else {
return Err(
"a body past the ceiling must be refused, not handed over as a prefix".into(),
);
};
assert_eq!(
error,
Error::BodyTooLarge {
limit: SMALL_CEILING
},
"the refusal names the ceiling that was reached"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_chunked_body_past_the_ceiling_is_refused() -> Result<(), Box<dyn std::error::Error>> {
const CHUNKED: &[u8] = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n40\r\nxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx\r\n40\r\nyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyyy\r\n0\r\n\r\n";
let (port, server) = serve_raw(CHUNKED.to_vec())?;
let Err(error) = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
) else {
return Err("a chunked body past the ceiling must be refused".into());
};
assert_eq!(
error,
Error::BodyTooLarge {
limit: SMALL_CEILING
},
"the ceiling is enforced against the decoded body, not the framing"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_close_delimited_body_past_the_ceiling_is_refused() -> Result<(), Box<dyn std::error::Error>>
{
let mut reply = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
reply.extend_from_slice(filler(SMALL_CEILING.saturating_mul(2)).as_bytes());
let (port, server) = serve_raw(reply)?;
let Err(error) = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
) else {
return Err("a close-delimited body past the ceiling must be refused".into());
};
assert_eq!(
error,
Error::BodyTooLarge {
limit: SMALL_CEILING
},
"a body with no declared length is bounded by the read, not by a header"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_preview_keeps_the_ceiling_and_reports_the_cut() -> Result<(), Box<dyn std::error::Error>> {
let payload = filler(SMALL_CEILING.saturating_mul(2));
let (port, server) = serve(vec![("200 OK", payload)])?;
let response = get_with(
&format!("http://127.0.0.1:{port}/"),
&previewing(SMALL_CEILING),
)?;
assert_eq!(
response.body.len(),
SMALL_CEILING,
"the preview stops at the ceiling and keeps no more"
);
assert!(
response.body.iter().all(|byte| *byte == b'x'),
"the preview is the body's own prefix"
);
assert_eq!(
response.truncation,
Truncation::Cut,
"a body that continues past the ceiling is reported as cut"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_preview_of_a_body_that_ends_under_the_ceiling_is_whole()
-> Result<(), Box<dyn std::error::Error>> {
let (port, server) = serve(vec![("200 OK", "alive".to_owned())])?;
let response = get_with(
&format!("http://127.0.0.1:{port}/"),
&previewing(SMALL_CEILING),
)?;
assert_eq!(response.body, b"alive");
assert_eq!(
response.truncation,
Truncation::Complete,
"a preview of a body that ended on its own is not a cut"
);
join_server(server)?;
Ok(())
}
#[test]
fn the_ceiling_is_enforced_before_any_utf8_decision() -> Result<(), Box<dyn std::error::Error>>
{
let mut reply = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
reply.extend_from_slice(&vec![0xFF_u8; SMALL_CEILING.saturating_mul(2)]);
let (port, server) = serve_raw(reply)?;
let Err(error) = get_with(
&format!("http://127.0.0.1:{port}/"),
&ceiling(SMALL_CEILING),
) else {
return Err("a multi-byte body past the ceiling must be refused".into());
};
assert_eq!(
error,
Error::BodyTooLarge {
limit: SMALL_CEILING
},
"the ceiling counts bytes, whatever they encode"
);
join_server(server)?;
Ok(())
}
#[test]
fn a_preview_stops_reading_at_the_ceiling() -> Result<(), Box<dyn std::error::Error>> {
let mut reply = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
reply.extend_from_slice(&vec![b'x'; SMALL_CEILING.saturating_mul(65_536)]);
let (port, server) = serve_raw(reply)?;
let response = get_with(
&format!("http://127.0.0.1:{port}/"),
&previewing(SMALL_CEILING),
)?;
assert_eq!(
response.body.len(),
SMALL_CEILING,
"the preview stops at the ceiling"
);
assert_eq!(
response.truncation,
Truncation::Cut,
"the reply continues past the ceiling"
);
let served = server
.join()
.map_err(|_| "the raw server thread panicked before replying")?;
assert!(
served.is_err(),
"a 4 MiB reply can only be written in full to a reader that keeps reading, so a \
completed write means the client drained it: a preview must stop at its ceiling"
);
Ok(())
}
}