#![cfg(not(target_arch = "wasm32"))]
use anyhow::Result;
use async_trait::async_trait;
use wacore::net::{HttpClient, HttpRequest, HttpResponse, StreamingHttpResponse, UploadBody};
use wacore::stats::HttpResourceReport;
pub const DEFAULT_MAX_BODY_BYTES: u64 = 2 * 1024 * 1024 * 1024;
const INPUT_BUFFER_BYTES: u64 = 16 * 1024;
const OUTPUT_BUFFER_BYTES: u64 = 16 * 1024;
const MAX_IDLE_CONNECTIONS: u64 = 3;
#[derive(Debug, Clone)]
pub struct UreqHttpClient {
agent: ureq::Agent,
max_body_bytes: u64,
pool_report: Option<HttpResourceReport>,
}
fn default_pool_report() -> HttpResourceReport {
HttpResourceReport {
pool_connections: Some(MAX_IDLE_CONNECTIONS),
pool_buffer_bytes: Some(MAX_IDLE_CONNECTIONS * (INPUT_BUFFER_BYTES + OUTPUT_BUFFER_BYTES)),
inflight_bytes: None,
}
}
impl UreqHttpClient {
pub fn new() -> Self {
Self {
agent: build_agent(),
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
pool_report: Some(default_pool_report()),
}
}
pub fn with_agent(agent: ureq::Agent) -> Self {
Self {
agent,
max_body_bytes: DEFAULT_MAX_BODY_BYTES,
pool_report: None,
}
}
pub fn with_max_body_bytes(mut self, max_body_bytes: u64) -> Self {
self.max_body_bytes = max_body_bytes;
self
}
}
impl Default for UreqHttpClient {
fn default() -> Self {
Self::new()
}
}
fn build_agent() -> ureq::Agent {
use ureq::config::Config;
#[allow(unused_mut)]
let mut builder = Config::builder()
.input_buffer_size(INPUT_BUFFER_BYTES as usize)
.output_buffer_size(OUTPUT_BUFFER_BYTES as usize)
.max_idle_connections(MAX_IDLE_CONNECTIONS as usize)
.max_idle_connections_per_host(2);
#[cfg(feature = "danger-skip-tls-verify")]
{
use ureq::tls::TlsConfig;
builder = builder.tls_config(TlsConfig::builder().disable_verification(true).build());
}
builder.build().into()
}
fn status_as_response<Any>(req: ureq::RequestBuilder<Any>) -> ureq::RequestBuilder<Any> {
req.config().http_status_as_error(false).build()
}
const ERROR_BODY_CAP: u64 = 64 * 1024;
fn read_body(response: ureq::http::Response<ureq::Body>, max_body_bytes: u64) -> Result<Vec<u8>> {
if response.status().is_success() {
return Ok(response
.into_body()
.into_with_config()
.limit(max_body_bytes)
.read_to_vec()?);
}
let mut body = Vec::new();
let mut reader = std::io::Read::take(
response.into_body().into_reader(),
max_body_bytes.min(ERROR_BODY_CAP),
);
let _ = std::io::Read::read_to_end(&mut reader, &mut body);
Ok(body)
}
#[async_trait]
impl HttpClient for UreqHttpClient {
async fn execute(&self, request: HttpRequest) -> Result<HttpResponse> {
let agent = self.agent.clone();
let max_body_bytes = self.max_body_bytes;
tokio::task::spawn_blocking(move || {
let response = match request.method.as_str() {
"GET" => {
let mut req = status_as_response(agent.get(&request.url));
for (key, value) in &request.headers {
req = req.header(key, value);
}
req.call()?
}
"POST" => {
let mut req = status_as_response(agent.post(&request.url));
for (key, value) in &request.headers {
req = req.header(key, value);
}
if let Some(body) = request.body {
req.send(&body[..])?
} else {
req.send(&[])?
}
}
method => {
return Err(anyhow::anyhow!("Unsupported HTTP method: {}", method));
}
};
let status_code = response.status().as_u16();
let body = read_body(response, max_body_bytes)?;
Ok(HttpResponse { status_code, body })
})
.await?
}
fn supports_streaming(&self) -> bool {
true
}
fn execute_streaming(&self, request: HttpRequest) -> Result<StreamingHttpResponse> {
let response = match request.method.as_str() {
"GET" => {
let mut req = status_as_response(self.agent.get(&request.url));
for (key, value) in &request.headers {
req = req.header(key, value);
}
req.call()?
}
method => {
return Err(anyhow::anyhow!(
"Streaming only supports GET, got: {}",
method
));
}
};
let status_code = response.status().as_u16();
let reader = std::io::Read::take(response.into_body().into_reader(), self.max_body_bytes);
Ok(StreamingHttpResponse {
status_code,
body: Box::new(reader),
})
}
fn supports_upload_streaming(&self) -> bool {
true
}
fn execute_upload(
&self,
request: HttpRequest,
body: UploadBody,
content_length: u64,
) -> Result<HttpResponse> {
if request.method != "POST" {
return Err(anyhow::anyhow!(
"Upload streaming only supports POST, got: {}",
request.method
));
}
let mut req = status_as_response(self.agent.post(&request.url));
for (key, value) in &request.headers {
req = req.header(key, value);
}
let content_length = content_length.to_string();
req = req.header("content-length", content_length.as_str());
let response = req.send(ureq::SendBody::from_owned_reader(body))?;
let status_code = response.status().as_u16();
let body = read_body(response, self.max_body_bytes)?;
Ok(HttpResponse { status_code, body })
}
fn resource_report(&self) -> Option<HttpResourceReport> {
self.pool_report
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
fn spawn_fixed_size_server(body_size: usize) -> String {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port");
let addr = listener.local_addr().unwrap();
thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept");
let mut buf = [0u8; 4096];
let mut total = Vec::new();
loop {
let n = stream.read(&mut buf).unwrap_or(0);
if n == 0 {
return;
}
total.extend_from_slice(&buf[..n]);
if total.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body_size
);
stream.write_all(header.as_bytes()).unwrap();
let chunk = vec![0xABu8; 64 * 1024];
let mut sent = 0usize;
while sent < body_size {
let take = chunk.len().min(body_size - sent);
stream.write_all(&chunk[..take]).unwrap();
sent += take;
}
});
format!("http://{}", addr)
}
#[tokio::test(flavor = "current_thread")]
async fn execute_accepts_body_larger_than_ureq_default_limit() {
const SIZE: usize = 12 * 1024 * 1024;
let url = spawn_fixed_size_server(SIZE);
let resp = UreqHttpClient::new()
.execute(HttpRequest {
method: "GET".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
})
.await
.expect("body must fit under the configured cap");
assert_eq!(resp.status_code, 200);
assert_eq!(resp.body.len(), SIZE);
}
#[tokio::test(flavor = "current_thread")]
async fn with_max_body_bytes_enforces_tighter_cap() {
const SIZE: usize = 4 * 1024 * 1024;
let url = spawn_fixed_size_server(SIZE);
UreqHttpClient::new()
.with_max_body_bytes(1024)
.execute(HttpRequest {
method: "GET".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
})
.await
.expect_err("1 KiB cap must reject a 4 MiB body");
}
#[tokio::test(flavor = "current_thread")]
async fn execute_streaming_bounds_body_at_cap() {
const SIZE: usize = 4 * 1024 * 1024;
const CAP: u64 = 1024;
let url = spawn_fixed_size_server(SIZE);
let read = tokio::task::spawn_blocking(move || {
let mut resp = UreqHttpClient::new()
.with_max_body_bytes(CAP)
.execute_streaming(HttpRequest {
method: "GET".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
})
.expect("streaming GET should start");
let mut sink = std::io::sink();
std::io::copy(&mut resp.body, &mut sink).expect("draining the reader should not error")
})
.await
.unwrap();
assert_eq!(read, CAP, "streaming body must stop at the cap");
}
fn spawn_capture_server() -> (String, std::sync::mpsc::Receiver<(String, Vec<u8>)>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port");
let addr = listener.local_addr().unwrap();
let (tx, rx) = std::sync::mpsc::channel();
thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept");
let mut buf = Vec::new();
let mut tmp = [0u8; 4096];
let header_end = loop {
if let Some(pos) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break pos + 4;
}
let n = stream.read(&mut tmp).unwrap_or(0);
if n == 0 {
return;
}
buf.extend_from_slice(&tmp[..n]);
};
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
let content_length = headers.lines().find_map(|l| {
let (k, v) = l.split_once(':')?;
if k.trim().eq_ignore_ascii_case("content-length") {
v.trim().parse::<usize>().ok()
} else {
None
}
});
let mut body = buf[header_end..].to_vec();
if let Some(cl) = content_length {
while body.len() < cl {
let n = stream.read(&mut tmp).unwrap_or(0);
if n == 0 {
break;
}
body.extend_from_slice(&tmp[..n]);
}
}
let _ = stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\n{}");
let _ = tx.send((headers, body));
});
(format!("http://{addr}"), rx)
}
fn parsed_content_length(headers: &str) -> Option<usize> {
headers.lines().find_map(|l| {
let (k, v) = l.split_once(':')?;
k.trim()
.eq_ignore_ascii_case("content-length")
.then(|| v.trim().parse::<usize>().ok())
.flatten()
})
}
#[test]
fn upload_streaming_sets_content_length_not_chunked() {
let (url, rx) = spawn_capture_server();
let payload: Vec<u8> = (0..5000u32).map(|i| i as u8).collect();
let client = UreqHttpClient::new();
let resp = client
.execute_upload(
HttpRequest {
method: "POST".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
},
Box::new(std::io::Cursor::new(payload.clone())),
payload.len() as u64,
)
.expect("upload should succeed");
assert_eq!(resp.status_code, 200);
let (headers, body) = rx
.recv_timeout(std::time::Duration::from_secs(5))
.expect("server should capture the request");
assert_eq!(
parsed_content_length(&headers),
Some(payload.len()),
"exact Content-Length expected, headers:\n{headers}"
);
assert!(
!headers.to_ascii_lowercase().contains("transfer-encoding"),
"body must not be chunked, headers:\n{headers}"
);
assert_eq!(body, payload, "server must receive the exact bytes");
}
#[test]
fn upload_streaming_large_body_integrity() {
let (url, rx) = spawn_capture_server();
let payload: Vec<u8> = (0..200_000usize).map(|i| (i % 251) as u8).collect();
let client = UreqHttpClient::new();
let resp = client
.execute_upload(
HttpRequest {
method: "POST".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
},
Box::new(std::io::Cursor::new(payload.clone())),
payload.len() as u64,
)
.expect("upload should succeed");
assert_eq!(resp.status_code, 200);
let (headers, body) = rx
.recv_timeout(std::time::Duration::from_secs(10))
.expect("server should capture the request");
assert_eq!(parsed_content_length(&headers), Some(payload.len()));
assert_eq!(body, payload);
}
#[test]
fn resource_report_estimates_default_pool() {
let report = UreqHttpClient::new()
.resource_report()
.expect("default agent reports a pool estimate");
assert_eq!(report.pool_connections, Some(MAX_IDLE_CONNECTIONS));
assert_eq!(
report.pool_buffer_bytes,
Some(MAX_IDLE_CONNECTIONS * (INPUT_BUFFER_BYTES + OUTPUT_BUFFER_BYTES))
);
assert_eq!(report.inflight_bytes, None);
assert!(report.total_bytes() > 0);
assert!(
UreqHttpClient::with_agent(build_agent())
.resource_report()
.is_none(),
"custom-agent client reports no estimate"
);
assert!(
UreqHttpClient::new()
.with_max_body_bytes(1024)
.resource_report()
.is_some()
);
}
fn spawn_status_server(status: u16, reason: &str) -> String {
spawn_status_server_with_body(status, reason, b"denied".to_vec())
}
fn spawn_status_server_with_body(status: u16, reason: &str, body: Vec<u8>) -> String {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port");
let addr = listener.local_addr().unwrap();
let reason = reason.to_string();
thread::spawn(move || {
let Ok((mut stream, _)) = listener.accept() else {
return;
};
let mut buf = Vec::new();
let mut tmp = [0u8; 4096];
let header_end = loop {
if let Some(pos) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break pos + 4;
}
match stream.read(&mut tmp) {
Ok(0) | Err(_) => return,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
};
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
if let Some(cl) = parsed_content_length(&headers) {
let mut body_len = buf.len() - header_end;
while body_len < cl {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => body_len += n,
}
}
}
let header = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = stream.write_all(header.as_bytes());
let _ = stream.write_all(&body);
});
format!("http://{addr}")
}
fn get(url: String) -> HttpRequest {
HttpRequest {
method: "GET".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
}
}
#[tokio::test(flavor = "current_thread")]
async fn execute_surfaces_non_2xx_status_instead_of_erroring() {
for (status, reason) in [
(401u16, "Unauthorized"),
(403, "Forbidden"),
(404, "Not Found"),
] {
let url = spawn_status_server(status, reason);
let resp = UreqHttpClient::new()
.execute(get(url))
.await
.unwrap_or_else(|e| panic!("{status} must arrive as a response, got error: {e}"));
assert_eq!(resp.status_code, status);
assert_eq!(resp.body, b"denied");
}
}
#[tokio::test(flavor = "current_thread")]
async fn execute_post_surfaces_non_2xx_status_instead_of_erroring() {
let url = spawn_status_server(403, "Forbidden");
let resp = UreqHttpClient::new()
.execute(HttpRequest::post(url).with_body(b"payload".to_vec()))
.await
.expect("403 must arrive as a response, not an error");
assert_eq!(resp.status_code, 403);
}
#[tokio::test(flavor = "current_thread")]
async fn execute_streaming_surfaces_non_2xx_status_instead_of_erroring() {
let url = spawn_status_server(403, "Forbidden");
let status = tokio::task::spawn_blocking(move || {
UreqHttpClient::new()
.execute_streaming(get(url))
.expect("403 must arrive as a response, not an error")
.status_code
})
.await
.unwrap();
assert_eq!(status, 403);
}
#[test]
fn execute_upload_surfaces_non_2xx_status_instead_of_erroring() {
let url = spawn_status_server(403, "Forbidden");
let payload = vec![7u8; 128];
let resp = UreqHttpClient::new()
.execute_upload(
HttpRequest {
method: "POST".into(),
url,
headers: std::collections::HashMap::new(),
body: None,
},
Box::new(std::io::Cursor::new(payload.clone())),
payload.len() as u64,
)
.expect("403 must arrive as a response, not an error");
assert_eq!(resp.status_code, 403);
}
#[tokio::test(flavor = "current_thread")]
async fn over_cap_error_body_does_not_cost_the_status() {
const CAP: u64 = 1024;
let url = spawn_status_server_with_body(403, "Forbidden", vec![b'x'; 4 * 1024 * 1024]);
let resp = UreqHttpClient::new()
.with_max_body_bytes(CAP)
.execute(get(url))
.await
.expect("an over-cap error page must not erase the status it came with");
assert_eq!(resp.status_code, 403);
assert!(
resp.body.len() as u64 <= CAP,
"the diagnostic body must stay bounded, got {} bytes",
resp.body.len()
);
}
#[tokio::test(flavor = "current_thread")]
async fn over_cap_success_body_is_still_an_error() {
let url = spawn_status_server_with_body(200, "OK", vec![b'x'; 4 * 1024 * 1024]);
UreqHttpClient::new()
.with_max_body_bytes(1024)
.execute(get(url))
.await
.expect_err("a truncated 2xx payload must never look like a complete one");
}
#[tokio::test(flavor = "current_thread")]
async fn custom_agent_also_surfaces_non_2xx_status() {
let url = spawn_status_server(403, "Forbidden");
let agent: ureq::Agent = ureq::config::Config::builder().build().into();
let resp = UreqHttpClient::with_agent(agent)
.execute(get(url))
.await
.expect("403 must arrive as a response even with a custom agent");
assert_eq!(resp.status_code, 403);
}
#[test]
fn upload_streaming_rejects_non_post() {
let client = UreqHttpClient::new();
let err = client.execute_upload(
HttpRequest {
method: "GET".into(),
url: "http://127.0.0.1:0/never".into(),
headers: std::collections::HashMap::new(),
body: None,
},
Box::new(std::io::Cursor::new(vec![1u8, 2, 3])),
3,
);
assert!(err.is_err(), "only POST is allowed for upload streaming");
}
}