use std::fmt;
use std::io::Read;
use std::time::Duration;
use iri_string::types::UriAbsoluteStr;
pub const DEFAULT_MAX_BODY_BYTES: usize = 8_388_608;
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)]
#[non_exhaustive]
pub struct Options {
pub timeout: Duration,
pub user_agent: String,
pub headers: Vec<(String, String)>,
pub max_body_bytes: usize,
pub body_policy: BodyPolicy,
}
impl Default for Options {
fn default() -> Self {
Self {
timeout: Duration::from_secs(30),
user_agent: format!("lgwks-std/{}", env!("CARGO_PKG_VERSION")),
headers: Vec::new(),
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
body_policy: BodyPolicy::Whole,
}
}
}
impl Options {
#[must_use]
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
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 idempotency_key(mut self, key: &str) -> Self {
self.headers
.push(("Idempotency-Key".to_owned(), key.to_owned()));
self
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Response {
pub status: u16,
pub headers: Vec<(String, String)>,
pub body: Vec<u8>,
pub truncation: Truncation,
}
impl Response {
pub fn text(&self) -> Result<&str, Error> {
std::str::from_utf8(&self.body).map_err(|utf8_error| {
Error::Transport(format!("response body is not valid UTF-8: {utf8_error}"))
})
}
#[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,
Timeout,
BodyTooLarge {
limit: usize,
},
Transport(String),
}
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::Timeout => write!(f, "request timed out"),
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::Transport(ref cause) => write!(f, "transport failure: {cause}"),
}
}
}
impl std::error::Error for Error {}
pub fn validate_url(url: &str) -> Result<(), Error> {
UriAbsoluteStr::new(url).map_err(|_malformed| Error::InvalidUrl)?;
let Some(scheme) = url.split_once(':').map(|(scheme, _)| scheme) else {
return Err(Error::InvalidUrl);
};
if !scheme.eq_ignore_ascii_case("http") && !scheme.eq_ignore_ascii_case("https") {
return Err(Error::InvalidUrl);
}
let Some(rest) = scheme
.len()
.checked_add(1)
.and_then(|after| url.get(after..))
else {
return Err(Error::InvalidUrl);
};
let Some(authority) = rest.strip_prefix("//") else {
return Err(Error::InvalidUrl);
};
let host_end = authority.find(['/', '?', '#']).unwrap_or(authority.len());
if authority[..host_end].is_empty() {
return Err(Error::InvalidUrl);
}
Ok(())
}
fn agent(options: &Options) -> ureq::Agent {
let config = ureq::Agent::config_builder()
.timeout_global(Some(options.timeout))
.http_status_as_error(false)
.user_agent(&options.user_agent)
.build();
ureq::Agent::new_with_config(config)
}
fn response_of(
mut response: ureq::http::Response<ureq::Body>,
options: &Options,
) -> Result<Response, Error> {
let status = response.status().as_u16();
let headers = response
.headers()
.iter()
.map(|(name, value)| {
(
name.to_string(),
String::from_utf8_lossy(value.as_bytes()).into_owned(),
)
})
.collect();
let (body, truncation) = read_bounded(&mut response.body_mut().as_reader(), options)?;
Ok(Response {
status,
headers,
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| Error::Transport(read_error.to_string()))?;
if read == 0 {
break Truncation::Complete;
}
body.extend_from_slice(&chunk[..read]);
};
if options.body_policy == BodyPolicy::Whole && truncation == Truncation::Cut {
return Err(Error::BodyTooLarge { limit: ceiling });
}
Ok((body, truncation))
}
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| Error::Transport(read_error.to_string()))?;
Ok(if read == 0 {
Truncation::Complete
} else {
Truncation::Cut
})
}
fn map_error(error: ureq::Error) -> Error {
match error {
ureq::Error::Timeout(_) => Error::Timeout,
ureq::Error::BadUri(_) => Error::InvalidUrl,
other => Error::Transport(other.to_string()),
}
}
pub fn get(url: &str) -> Result<Response, Error> {
get_with(url, &Options::default())
}
pub fn get_with(url: &str, options: &Options) -> Result<Response, Error> {
validate_url(url)?;
let mut call = agent(options).get(url);
for header in &options.headers {
call = call.header(header.0.as_str(), header.1.as_str());
}
call.call()
.map_err(map_error)
.and_then(|response| response_of(response, 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> {
validate_url(url)?;
let mut request = agent(options)
.post(url)
.header("Content-Type", content_type);
for header in &options.headers {
request = request.header(header.0.as_str(), header.1.as_str());
}
let response = request.send(body).map_err(map_error)?;
response_of(response, options)
}
#[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;
use std::thread;
const ECHO: &str = "echo-body-123";
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()?;
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())?;
}
Ok(())
});
Ok((port, handle))
}
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),
user_agent: "lgwks-std-test".into(),
headers: Vec::new(),
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
body_policy: BodyPolicy::Whole,
}
}
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 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::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 mut options = quiet();
options
.headers
.push(("Authorization".into(), "Bearer test-token".into()));
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}"
);
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 handle = thread::spawn(move || -> std::io::Result<()> {
let (_stream, _) = listener.accept()?;
thread::sleep(Duration::from_secs(30));
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_eq!(error, Error::Timeout);
drop(handle);
Ok(())
}
#[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 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(())
}
}