use std::{borrow::Cow, collections::HashMap, str::FromStr, sync::Arc, time::Duration};
use nautilus_core::collections::into_ustr_vec;
use nautilus_cryptography::providers::install_cryptographic_provider;
use reqwest::{
Method, Response, Url,
header::{HeaderMap, HeaderName, HeaderValue},
};
use ustr::Ustr;
use super::{HttpClientError, HttpResponse, HttpStatus};
use crate::ratelimiter::{RateLimiter, clock::MonotonicClock, quota::Quota};
const DEFAULT_POOL_MAX_IDLE_PER_HOST: usize = 32;
const DEFAULT_POOL_IDLE_TIMEOUT_SECS: u64 = 60;
const DEFAULT_HTTP2_KEEP_ALIVE_SECS: u64 = 30;
const DEFAULT_MAX_RESPONSE_BYTES: usize = 100 * 1024 * 1024;
#[derive(Clone, Debug)]
pub struct HttpClient {
pub(crate) client: InnerHttpClient,
pub(crate) rate_limiters: Arc<[Arc<RateLimiter<Ustr, MonotonicClock>>]>,
}
impl HttpClient {
pub fn new(
headers: HashMap<String, String>,
header_keys: Vec<String>,
keyed_quotas: Vec<(String, Quota)>,
default_quota: Option<Quota>,
timeout_secs: Option<u64>,
proxy_url: Option<String>,
) -> Result<Self, HttpClientError> {
let keyed_quotas = keyed_quotas
.into_iter()
.map(|(key, quota)| (Ustr::from(&key), quota))
.collect();
let rate_limiter = Arc::new(RateLimiter::new_with_quota(default_quota, keyed_quotas));
Self::new_with_rate_limiter(headers, header_keys, timeout_secs, proxy_url, rate_limiter)
}
pub fn new_with_rate_limiter(
headers: HashMap<String, String>,
header_keys: Vec<String>,
timeout_secs: Option<u64>,
proxy_url: Option<String>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
) -> Result<Self, HttpClientError> {
Self::new_with_rate_limiters(
headers,
header_keys,
timeout_secs,
proxy_url,
vec![rate_limiter],
)
}
pub fn new_with_rate_limiters(
headers: HashMap<String, String>,
header_keys: Vec<String>,
timeout_secs: Option<u64>,
proxy_url: Option<String>,
rate_limiters: Vec<Arc<RateLimiter<Ustr, MonotonicClock>>>,
) -> Result<Self, HttpClientError> {
install_cryptographic_provider();
let mut header_map = HeaderMap::new();
for (key, value) in headers {
let header_name = HeaderName::from_str(&key)
.map_err(|e| HttpClientError::Error(format!("Invalid header name '{key}': {e}")))?;
let header_value = HeaderValue::from_str(&value).map_err(|e| {
HttpClientError::Error(format!("Invalid header value for '{key}': {e}"))
})?;
header_map.insert(header_name, header_value);
}
let mut client_builder = reqwest::Client::builder()
.default_headers(header_map)
.tcp_nodelay(true)
.pool_max_idle_per_host(DEFAULT_POOL_MAX_IDLE_PER_HOST)
.pool_idle_timeout(Duration::from_secs(DEFAULT_POOL_IDLE_TIMEOUT_SECS))
.http2_keep_alive_interval(Duration::from_secs(DEFAULT_HTTP2_KEEP_ALIVE_SECS))
.http2_keep_alive_while_idle(true)
.http2_adaptive_window(true);
if let Some(timeout_secs) = timeout_secs {
client_builder = client_builder.timeout(Duration::from_secs(timeout_secs));
}
if let Some(proxy_url) = proxy_url {
let proxy = reqwest::Proxy::all(&proxy_url)
.map_err(|_| HttpClientError::InvalidProxy("proxy URL is malformed".to_string()))?;
client_builder = client_builder.proxy(proxy);
}
let client = client_builder
.build()
.map_err(|e| HttpClientError::ClientBuildError(e.to_string()))?;
let response_headers = header_keys
.into_iter()
.map(|key| match HeaderName::from_str(&key) {
Ok(name) => Ok((key, name)),
Err(e) => Err(HttpClientError::Error(format!(
"Invalid header key '{key}': {e}"
))),
})
.collect::<Result<Vec<_>, _>>()?;
let client = InnerHttpClient {
client,
response_headers: Arc::from(response_headers),
max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
};
Ok(Self {
client,
rate_limiters: rate_limiters.into(),
})
}
#[expect(clippy::too_many_arguments)]
pub async fn request(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
let keys = keys.map(into_ustr_vec);
self.request_with_ustr_keys(method, url, params, headers, body, timeout_secs, keys)
.await
}
#[expect(clippy::too_many_arguments)]
pub async fn request_with_url_redacted(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
let keys = keys.map(into_ustr_vec);
self.await_rate_limits(keys.as_deref()).await;
self.client
.send_request_with_url_redacted(method, url, params, headers, body, timeout_secs)
.await
}
#[expect(clippy::too_many_arguments)]
pub async fn request_with_params<P: serde::Serialize>(
&self,
method: Method,
url: String,
params: Option<&P>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
let keys = keys.map(into_ustr_vec);
self.await_rate_limits(keys.as_deref()).await;
self.client
.send_request_with_query(method, url, params, headers, body, timeout_secs)
.await
}
#[expect(clippy::too_many_arguments)]
pub async fn request_with_ustr_keys(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<Ustr>>,
) -> Result<HttpResponse, HttpClientError> {
self.await_rate_limits(keys.as_deref()).await;
self.client
.send_request(method, url, params, headers, body, timeout_secs)
.await
}
pub(crate) async fn await_rate_limits(&self, keys: Option<&[Ustr]>) {
for rate_limiter in self.rate_limiters.iter() {
rate_limiter.await_keys_ready(keys).await;
}
}
pub async fn get(
&self,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
self.request(Method::GET, url, params, headers, None, timeout_secs, keys)
.await
}
pub async fn post(
&self,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
self.request(Method::POST, url, params, headers, body, timeout_secs, keys)
.await
}
pub async fn patch(
&self,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
self.request(
Method::PATCH,
url,
params,
headers,
body,
timeout_secs,
keys,
)
.await
}
pub async fn delete(
&self,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
timeout_secs: Option<u64>,
keys: Option<Vec<String>>,
) -> Result<HttpResponse, HttpClientError> {
self.request(
Method::DELETE,
url,
params,
headers,
None,
timeout_secs,
keys,
)
.await
}
}
#[derive(Clone, Debug)]
pub struct InnerHttpClient {
pub(crate) client: reqwest::Client,
pub(crate) response_headers: Arc<[(String, HeaderName)]>,
pub(crate) max_response_bytes: usize,
}
impl InnerHttpClient {
pub async fn send_request(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
) -> Result<HttpResponse, HttpClientError> {
self.send_request_with_redaction(method, url, params, headers, body, timeout_secs, false)
.await
}
async fn send_request_with_url_redacted(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
) -> Result<HttpResponse, HttpClientError> {
self.send_request_with_redaction(method, url, params, headers, body, timeout_secs, true)
.await
}
#[expect(clippy::too_many_arguments)]
async fn send_request_with_redaction(
&self,
method: Method,
url: String,
params: Option<&HashMap<String, Vec<String>>>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
redact_url: bool,
) -> Result<HttpResponse, HttpClientError> {
let full_url = encode_url_params(&url, params)?;
self.send_request_internal(
method,
full_url.as_ref(),
None::<&()>,
headers,
body,
timeout_secs,
redact_url,
)
.await
}
pub async fn send_request_with_query<Q: serde::Serialize>(
&self,
method: Method,
url: String,
query: Option<&Q>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
) -> Result<HttpResponse, HttpClientError> {
self.send_request_internal(method, &url, query, headers, body, timeout_secs, false)
.await
}
#[expect(clippy::too_many_arguments)]
async fn send_request_internal<Q: serde::Serialize>(
&self,
method: Method,
url: &str,
query: Option<&Q>,
headers: Option<HashMap<String, String>>,
body: Option<Vec<u8>>,
timeout_secs: Option<u64>,
redact_url: bool,
) -> Result<HttpResponse, HttpClientError> {
let reqwest_url =
Url::parse(url).map_err(|e| HttpClientError::from(format!("URL parse error: {e}")))?;
let mut request_builder = self.client.request(method, reqwest_url);
let extra_header_count = headers.as_ref().map_or(0, HashMap::len);
let body_len = body.as_ref().map_or(0, Vec::len);
if let Some(headers) = headers {
let mut header_map = HeaderMap::with_capacity(headers.len());
for (header_key, header_value) in &headers {
let key = HeaderName::from_bytes(header_key.as_bytes())
.map_err(|e| HttpClientError::from(format!("Invalid header name: {e}")))?;
if header_map
.insert(
key.clone(),
header_value.parse().map_err(|e| {
HttpClientError::from(format!("Invalid header value: {e}"))
})?,
)
.is_some()
{
log::trace!("Replaced duplicate request header '{key}'");
}
}
request_builder = request_builder.headers(header_map);
}
if let Some(q) = query {
request_builder = request_builder.query(q);
}
if let Some(timeout_secs) = timeout_secs {
request_builder = request_builder.timeout(Duration::new(timeout_secs, 0));
}
let request = match body {
Some(b) => request_builder
.body(b)
.build()
.map_err(|e| http_client_error(e, redact_url))?,
None => request_builder
.build()
.map_err(|e| http_client_error(e, redact_url))?,
};
let query_len = request.url().query().map_or(0, str::len);
log::trace!(
"Sending HTTP request: method={} extra_headers={extra_header_count} \
query_bytes={query_len} body_bytes={body_len}",
request.method(),
);
let response = self
.client
.execute(request)
.await
.map_err(|e| http_client_error(e, redact_url))?;
self.to_response_internal(response, redact_url).await
}
pub async fn to_response(&self, response: Response) -> Result<HttpResponse, HttpClientError> {
self.to_response_internal(response, false).await
}
async fn to_response_internal(
&self,
response: Response,
redact_url: bool,
) -> Result<HttpResponse, HttpClientError> {
let status_code = response.status();
let resp_headers = response.headers();
let header_count = resp_headers.len();
let mut headers = HashMap::with_capacity(std::cmp::min(
self.response_headers.len(),
resp_headers.len(),
));
for (key, name) in self.response_headers.iter() {
if let Some(val) = resp_headers.get(name)
&& let Ok(v) = val.to_str()
{
headers.insert(key.clone(), v.to_owned());
}
}
let status = HttpStatus::new(status_code);
let body = self.read_body_capped(response, redact_url).await?;
log::trace!(
"Received HTTP response: status={status_code} headers={header_count} body_bytes={}",
body.len(),
);
Ok(HttpResponse {
status,
headers,
body,
})
}
async fn read_body_capped(
&self,
mut response: Response,
redact_url: bool,
) -> Result<bytes::Bytes, HttpClientError> {
let max = self.max_response_bytes;
if let Some(len) = response.content_length()
&& len > max as u64
{
return Err(HttpClientError::Error(format!(
"HTTP response body of {len} bytes exceeds maximum of {max} bytes",
)));
}
let mut buf = bytes::BytesMut::new();
while let Some(chunk) = response
.chunk()
.await
.map_err(|e| http_client_error(e, redact_url))?
{
if buf.len() + chunk.len() > max {
return Err(HttpClientError::Error(format!(
"HTTP response body exceeds maximum of {max} bytes",
)));
}
buf.extend_from_slice(&chunk);
}
Ok(buf.freeze())
}
}
fn http_client_error(error: reqwest::Error, redact_url: bool) -> HttpClientError {
if redact_url {
HttpClientError::from(error.without_url())
} else {
HttpClientError::from(error)
}
}
impl Default for InnerHttpClient {
fn default() -> Self {
install_cryptographic_provider();
let client = reqwest::Client::new();
Self {
client,
response_headers: Arc::default(),
max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
}
}
}
fn encode_url_params<'a>(
url: &'a str,
params: Option<&HashMap<String, Vec<String>>>,
) -> Result<Cow<'a, str>, HttpClientError> {
let Some(params) = params else {
return Ok(Cow::Borrowed(url));
};
let pairs: Vec<(&str, &str)> = params
.iter()
.flat_map(|(key, values)| {
values
.iter()
.map(move |value| (key.as_str(), value.as_str()))
})
.collect();
if pairs.is_empty() {
return Ok(Cow::Borrowed(url));
}
let query_string = serde_urlencoded::to_string(pairs)
.map_err(|e| HttpClientError::Error(format!("Failed to encode params: {e}")))?;
let (base, fragment) = match url.split_once('#') {
Some((base, fragment)) => (base, Some(fragment)),
None => (url, None),
};
let separator = if base.contains('?') { '&' } else { '?' };
Ok(Cow::Owned(match fragment {
Some(fragment) => format!("{base}{separator}{query_string}#{fragment}"),
None => format!("{base}{separator}{query_string}"),
}))
}
#[cfg(test)]
mod encode_url_params_tests {
use std::{borrow::Cow, collections::HashMap};
use rstest::rstest;
use super::encode_url_params;
fn params(pairs: &[(&str, &str)]) -> HashMap<String, Vec<String>> {
let mut map: HashMap<String, Vec<String>> = HashMap::new();
for (key, value) in pairs {
map.entry((*key).to_string())
.or_default()
.push((*value).to_string());
}
map
}
#[rstest]
#[case("https://x/y", "https://x/y?a=b")]
#[case("https://x/y?old=1", "https://x/y?old=1&a=b")]
#[case("https://x/y#frag", "https://x/y?a=b#frag")]
#[case("https://x/y?old=1#frag", "https://x/y?old=1&a=b#frag")]
#[case(
"https://x/y#section?display=full",
"https://x/y?a=b#section?display=full"
)]
#[case("https://x/y#", "https://x/y?a=b#")]
fn test_query_is_inserted_before_the_fragment(#[case] url: &str, #[case] expected: &str) {
let params = params(&[("a", "b")]);
assert_eq!(encode_url_params(url, Some(¶ms)).unwrap(), expected);
}
#[rstest]
fn test_url_is_borrowed_when_no_params_are_supplied() {
assert!(matches!(
encode_url_params("https://x/y#frag", None).unwrap(),
Cow::Borrowed("https://x/y#frag")
));
}
#[rstest]
fn test_url_is_borrowed_when_params_are_empty() {
let params = HashMap::new();
assert!(matches!(
encode_url_params("https://x/y#frag", Some(¶ms)).unwrap(),
Cow::Borrowed("https://x/y#frag")
));
}
}
#[cfg(test)]
#[cfg(target_os = "linux")] mod tests {
use std::{net::SocketAddr, num::NonZeroU32, sync::Mutex};
use axum::{
Router,
body::to_bytes,
extract::Request,
response::IntoResponse,
routing::{any, delete, get, patch, post},
serve,
};
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
use http::status::StatusCode;
use log::{Level, LevelFilter, Log, Metadata, Record};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
sync::oneshot,
};
use super::*;
async fn capture_request(request: Request) -> impl IntoResponse {
let (parts, body) = request.into_parts();
let body = to_bytes(body, usize::MAX).await.unwrap();
let default_header = parts.headers.get("x-default").unwrap().to_str().unwrap();
let request_header = parts.headers.get("x-request").unwrap().to_str().unwrap();
let query = parts.uri.query().unwrap_or_default();
let body = String::from_utf8(body.to_vec()).unwrap();
let capture = format!(
"{}\n{}\n{query}\n{default_header}\n{request_header}\n{body}",
parts.method,
parts.uri.path(),
);
([("x-response-id", "response-42")], capture)
}
#[derive(Default)]
struct CapturingTraceLogger {
messages: Mutex<Vec<(Level, String)>>,
}
impl CapturingTraceLogger {
fn clear(&self) {
self.messages.lock().unwrap().clear();
}
fn messages(&self) -> Vec<(Level, String)> {
self.messages.lock().unwrap().clone()
}
}
impl Log for CapturingTraceLogger {
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
metadata.level() <= Level::Trace
}
fn log(&self, record: &Record<'_>) {
if self.enabled(record.metadata()) {
self.messages
.lock()
.unwrap()
.push((record.level(), record.args().to_string()));
}
}
fn flush(&self) {}
}
static CAPTURING_TRACE_LOGGER: CapturingTraceLogger = CapturingTraceLogger {
messages: Mutex::new(Vec::new()),
};
fn create_router() -> Router {
Router::new()
.route("/get", get(|| async { "hello-world!" }))
.route("/post", post(|| async { StatusCode::OK }))
.route("/patch", patch(|| async { StatusCode::OK }))
.route("/delete", delete(|| async { StatusCode::OK }))
.route("/capture", any(capture_request))
.route("/notfound", get(|| async { StatusCode::NOT_FOUND }))
.route(
"/slow",
get(|| async {
tokio::time::sleep(Duration::from_secs(2)).await;
"Eventually responded"
}),
)
.route(
"/large",
get(|| async { "x".repeat(1024 * 1024) }),
)
}
async fn start_test_server() -> Result<SocketAddr, Box<dyn std::error::Error + Send + Sync>> {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
serve(listener, create_router()).await.unwrap();
});
Ok(addr)
}
async fn spawn_rejecting_connect_proxy() -> (SocketAddr, oneshot::Receiver<String>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (request_tx, request_rx) = oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = Vec::new();
let mut chunk = [0u8; 1024];
loop {
let read = stream.read(&mut chunk).await.unwrap();
if read == 0 {
break;
}
request.extend_from_slice(&chunk[..read]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
request_tx
.send(String::from_utf8(request).unwrap())
.unwrap();
stream
.write_all(
b"HTTP/1.1 407 Proxy Authentication Required\r\nContent-Length: 0\r\n\r\n",
)
.await
.unwrap();
});
(addr, request_rx)
}
#[tokio::test]
async fn test_http_client_awaits_multiple_rate_limiters() {
let quota = Quota::per_minute(NonZeroU32::MIN);
let request_key = Ustr::from("scope:request");
let order_key = Ustr::from("scope:order");
let request_limiter = Arc::new(RateLimiter::new_with_quota(
None,
vec![(request_key, quota)],
));
let order_limiter = Arc::new(RateLimiter::new_with_quota(None, vec![(order_key, quota)]));
let client = HttpClient::new_with_rate_limiters(
HashMap::new(),
Vec::new(),
None,
None,
vec![Arc::clone(&request_limiter), Arc::clone(&order_limiter)],
)
.unwrap();
client
.await_rate_limits(Some(&[request_key, order_key]))
.await;
assert!(request_limiter.check_key(&request_key).is_err());
assert!(order_limiter.check_key(&order_key).is_err());
}
#[tokio::test]
async fn test_get() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient::default();
let response = client
.send_request(
reqwest::Method::GET,
format!("{url}/get"),
None,
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(String::from_utf8_lossy(&response.body), "hello-world!");
}
#[tokio::test]
async fn test_request_preserves_wire_semantics_and_extracts_response_headers() {
let addr = start_test_server().await.unwrap();
let mut default_headers = HashMap::new();
default_headers.insert("x-default".to_string(), "default-a".to_string());
let client = HttpClient::new(
default_headers,
vec!["x-response-id".to_string()],
vec![],
None,
None,
None,
)
.unwrap();
let mut params = HashMap::new();
params.insert(
"tag".to_string(),
vec!["A B".to_string(), "C/D".to_string()],
);
let mut request_headers = HashMap::new();
request_headers.insert("x-request".to_string(), "request-b".to_string());
let response = client
.request(
Method::PUT,
format!("http://{addr}/capture?existing=seed"),
Some(¶ms),
Some(request_headers),
Some(b"payload-c".to_vec()),
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(
response.headers,
HashMap::from([("x-response-id".to_string(), "response-42".to_string())])
);
assert_eq!(
response.body.as_ref(),
b"PUT\n/capture\nexisting=seed&tag=A+B&tag=C%2FD\ndefault-a\nrequest-b\npayload-c"
);
}
#[tokio::test]
async fn test_request_with_params_serializes_query_fields() {
#[derive(serde::Serialize)]
struct Query<'a> {
symbol: &'a str,
limit: u32,
}
let addr = start_test_server().await.unwrap();
let mut default_headers = HashMap::new();
default_headers.insert("x-default".to_string(), "default-d".to_string());
let client = HttpClient::new(
default_headers,
vec!["x-response-id".to_string()],
vec![],
None,
None,
None,
)
.unwrap();
let mut request_headers = HashMap::new();
request_headers.insert("x-request".to_string(), "request-e".to_string());
let params = Query {
symbol: "BTC/USDT",
limit: 37,
};
let response = client
.request_with_params(
Method::GET,
format!("http://{addr}/capture"),
Some(¶ms),
Some(request_headers),
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(
response.headers,
HashMap::from([("x-response-id".to_string(), "response-42".to_string())])
);
assert_eq!(
response.body.as_ref(),
b"GET\n/capture\nsymbol=BTC%2FUSDT&limit=37\ndefault-d\nrequest-e\n"
);
}
#[tokio::test]
async fn test_response_body_within_cap_is_returned() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient {
max_response_bytes: 4 * 1024 * 1024,
..Default::default()
};
let response = client
.send_request(
reqwest::Method::GET,
format!("{url}/large"),
None,
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(response.body.len(), 1024 * 1024);
}
#[tokio::test]
async fn test_response_body_exceeding_cap_is_rejected() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient {
max_response_bytes: 16 * 1024,
..Default::default()
};
let result = client
.send_request(
reqwest::Method::GET,
format!("{url}/large"),
None,
None,
None,
None,
)
.await;
let err = result.expect_err("oversized response body should be rejected");
assert!(
err.to_string().contains("exceeds maximum"),
"unexpected error: {err}",
);
}
#[tokio::test]
async fn test_post() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient::default();
let response = client
.send_request(
reqwest::Method::POST,
format!("{url}/post"),
None,
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_post_with_body() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient::default();
let mut body = HashMap::new();
body.insert(
"key1".to_string(),
serde_json::Value::String("value1".to_string()),
);
body.insert(
"key2".to_string(),
serde_json::Value::String("value2".to_string()),
);
let body_string = serde_json::to_string(&body).unwrap();
let body_bytes = body_string.into_bytes();
let response = client
.send_request(
reqwest::Method::POST,
format!("{url}/post"),
None,
None,
Some(body_bytes),
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_patch() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient::default();
let response = client
.send_request(
reqwest::Method::PATCH,
format!("{url}/patch"),
None,
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_delete() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}");
let client = InnerHttpClient::default();
let response = client
.send_request(
reqwest::Method::DELETE,
format!("{url}/delete"),
None,
None,
None,
None,
)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_not_found() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/notfound");
let client = InnerHttpClient::default();
let response = client
.send_request(reqwest::Method::GET, url, None, None, None, None)
.await
.unwrap();
assert!(response.status.is_client_error());
assert_eq!(response.status.as_u16(), 404);
}
#[tokio::test]
async fn test_timeout() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/slow");
let client = InnerHttpClient::default();
let result = client
.send_request(reqwest::Method::GET, url, None, None, None, Some(1))
.await;
assert!(
matches!(&result, Err(HttpClientError::TimeoutError(_))),
"Expected a timeout error, was: {result:?}"
);
}
#[rstest]
fn test_http_client_without_proxy() {
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
None, );
assert!(result.is_ok());
}
#[tokio::test]
async fn test_http_client_without_proxy_requests_directly() {
let addr = start_test_server().await.unwrap();
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, Some(2), None).unwrap();
let response = client
.request(
Method::GET,
format!("http://{addr}/get"),
None,
None,
None,
None,
None,
)
.await
.expect("direct request");
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(response.body.as_ref(), b"hello-world!");
}
#[tokio::test]
async fn test_http_client_redacted_url_request_preserves_response() {
let addr = start_test_server().await.unwrap();
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, Some(2), None).unwrap();
let response = client
.request_with_url_redacted(
Method::GET,
format!("http://{addr}/get"),
None,
None,
None,
None,
None,
)
.await
.expect("direct request with URL redaction");
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(response.body.as_ref(), b"hello-world!");
}
#[tokio::test]
async fn test_http_client_redacted_url_request_removes_endpoint_from_error() {
const USERINFO_SECRET: &str = "transport-userinfo-secret";
const PATH_SECRET: &str = "transport-path-secret";
const QUERY_SECRET: &str = "transport-query-secret";
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
drop(listener);
let url = format!(
"http://rpc-user:{USERINFO_SECRET}@{addr}/{PATH_SECRET}?api_key={QUERY_SECRET}"
);
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, Some(1), None).unwrap();
let error = client
.request_with_url_redacted(Method::GET, url.clone(), None, None, None, None, None)
.await
.expect_err("an unreachable endpoint should fail");
for rendered in [error.to_string(), format!("{error:?}")] {
assert!(!rendered.contains(USERINFO_SECRET));
assert!(!rendered.contains(PATH_SECRET));
assert!(!rendered.contains(QUERY_SECRET));
assert!(!rendered.contains(&url));
}
}
#[tokio::test]
async fn test_http_client_redacted_url_request_removes_endpoint_from_trace_logs() {
const USERINFO_SECRET: &str = "trace-userinfo-secret";
const PATH_SECRET: &str = "trace-path-secret";
const QUERY_SECRET: &str = "trace-query-secret";
let _ = log::set_logger(&CAPTURING_TRACE_LOGGER);
log::set_max_level(LevelFilter::Trace);
CAPTURING_TRACE_LOGGER.clear();
let addr = start_test_server().await.unwrap();
let url = format!(
"http://rpc-user:{USERINFO_SECRET}@{addr}/{PATH_SECRET}?api_key={QUERY_SECRET}"
);
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, Some(2), None).unwrap();
let response = client
.request_with_url_redacted(Method::GET, url.clone(), None, None, None, None, None)
.await
.expect("credentialized endpoint should return an HTTP response");
let messages = CAPTURING_TRACE_LOGGER.messages();
assert_eq!(response.status.as_u16(), StatusCode::NOT_FOUND.as_u16());
assert!(messages.iter().any(|(level, message)| {
*level == Level::Trace && message.starts_with("Sending HTTP request: method=GET")
}));
assert!(messages.iter().any(|(level, message)| {
*level == Level::Trace
&& message.starts_with("Received HTTP response: status=404 Not Found")
}));
for (_, message) in messages {
assert!(!message.contains(USERINFO_SECRET));
assert!(!message.contains(PATH_SECRET));
assert!(!message.contains(QUERY_SECRET));
assert!(!message.contains(&url));
}
}
#[tokio::test]
async fn test_http_client_uses_connect_and_proxy_authorization_for_https() {
const USERNAME: &str = "proxytest";
const PASSWORD: &str = "fixture42";
let (proxy_addr, request_rx) = spawn_rejecting_connect_proxy().await;
let client = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
Some(2),
Some(format!("http://{USERNAME}:{PASSWORD}@{proxy_addr}")),
)
.unwrap();
let error = client
.request(
Method::GET,
"https://fixture.example.test/path".to_string(),
None,
None,
None,
None,
None,
)
.await
.expect_err("proxy should reject CONNECT");
let request = request_rx.await.expect("captured CONNECT request");
let mut lines = request.split("\r\n");
let request_line = lines.next().expect("CONNECT request line");
let auth_value = lines
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("proxy-authorization")
.then_some(value.trim())
})
.expect("Proxy-Authorization header");
let expected_auth = format!("Basic {}", BASE64.encode(format!("{USERNAME}:{PASSWORD}")));
assert_eq!(request_line, "CONNECT fixture.example.test:443 HTTP/1.1");
assert_eq!(auth_value, expected_auth);
assert!(!error.to_string().contains(PASSWORD));
assert!(!error.to_string().contains(&BASE64.encode(PASSWORD)));
assert!(!error.to_string().contains(&expected_auth));
}
#[tokio::test]
async fn test_http_client_unreachable_proxy_error_redacts_credentials() {
const USERNAME: &str = "proxy-user";
const SECRET: &str = "unreachable-proxy-secret";
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = listener.local_addr().unwrap();
drop(listener);
let client = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
Some(1),
Some(format!("http://{USERNAME}:{SECRET}@{proxy_addr}")),
)
.unwrap();
let error = client
.request(
Method::GET,
"https://fixture.example.test/".to_string(),
None,
None,
None,
None,
None,
)
.await
.expect_err("unreachable proxy should fail");
assert!(!error.to_string().contains(SECRET));
assert!(!error.to_string().contains(&BASE64.encode(SECRET)));
assert!(
!error
.to_string()
.contains(&BASE64.encode(format!("{USERNAME}:{SECRET}")))
);
}
#[rstest]
fn test_http_client_with_valid_proxy() {
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
Some("http://proxy.example.com:8080".to_string()),
);
assert!(result.is_ok());
}
#[rstest]
fn test_http_client_with_socks5_proxy() {
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
Some("socks5://127.0.0.1:1080".to_string()),
);
assert!(result.is_ok());
}
#[rstest]
fn test_http_client_with_malformed_proxy() {
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
Some("://invalid".to_string()),
);
assert!(result.is_err());
assert!(matches!(result, Err(HttpClientError::InvalidProxy(_))));
}
#[rstest]
fn test_http_client_invalid_proxy_error_redacts_credentials() {
const SECRET: &str = "unique-proxy-secret";
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
Some(format!("http://proxytest:{SECRET}@[::1")),
);
let error = result.expect_err("malformed proxy URL should fail");
assert_eq!(
error.to_string(),
"Invalid proxy URL: proxy URL is malformed"
);
assert!(!error.to_string().contains(SECRET));
}
#[rstest]
fn test_http_client_with_empty_proxy_string() {
let result = HttpClient::new(
HashMap::new(),
vec![],
vec![],
None,
None,
Some(String::new()),
);
assert!(result.is_err());
assert!(matches!(result, Err(HttpClientError::InvalidProxy(_))));
}
#[tokio::test]
async fn test_http_client_get() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/get");
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, None, None).unwrap();
let response = client.get(url, None, None, None, None).await.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
assert_eq!(String::from_utf8_lossy(&response.body), "hello-world!");
}
#[tokio::test]
async fn test_http_client_post() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/post");
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, None, None).unwrap();
let response = client
.post(url, None, None, None, None, None)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_http_client_patch() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/patch");
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, None, None).unwrap();
let response = client
.patch(url, None, None, None, None, None)
.await
.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
#[tokio::test]
async fn test_http_client_delete() {
let addr = start_test_server().await.unwrap();
let url = format!("http://{addr}/delete");
let client = HttpClient::new(HashMap::new(), vec![], vec![], None, None, None).unwrap();
let response = client.delete(url, None, None, None, None).await.unwrap();
assert_eq!(response.status.as_u16(), StatusCode::OK.as_u16());
}
}