use async_trait::async_trait;
use std::fmt::Debug;
#[doc(no_inline)]
pub use bytes::Bytes;
#[doc(no_inline)]
pub use http::{Request, Response};
use opentelemetry::propagation::{Extractor, Injector};
pub struct HeaderInjector<'a>(pub &'a mut http::HeaderMap);
impl Injector for HeaderInjector<'_> {
fn set(&mut self, key: &str, value: String) {
if let Ok(name) = http::header::HeaderName::from_bytes(key.as_bytes()) {
if let Ok(val) = http::header::HeaderValue::from_str(&value) {
self.0.insert(name, val);
}
}
}
fn reserve(&mut self, additional: usize) {
self.0.reserve(additional);
}
}
pub struct HeaderExtractor<'a>(pub &'a http::HeaderMap);
impl Extractor for HeaderExtractor<'_> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).and_then(|value| value.to_str().ok())
}
fn keys(&self) -> Vec<&str> {
self.0
.keys()
.map(|value| value.as_str())
.collect::<Vec<_>>()
}
fn get_all(&self, key: &str) -> Option<Vec<&str>> {
let all_iter = self.0.get_all(key).iter();
if let (0, Some(0)) = all_iter.size_hint() {
return None;
}
Some(all_iter.filter_map(|value| value.to_str().ok()).collect())
}
}
pub type HttpError = Box<dyn std::error::Error + Send + Sync + 'static>;
#[async_trait]
pub trait HttpClient: Debug + Send + Sync {
async fn send_bytes(&self, request: Request<Bytes>) -> Result<Response<Bytes>, HttpError>;
}
#[cfg(any(feature = "reqwest", feature = "hyper"))]
const MAX_RESPONSE_BODY_BYTES: usize = 4 * 1024 * 1024;
#[derive(Debug, Default)]
pub struct ResponseBodyTooLarge {
_private: (),
}
impl ResponseBodyTooLarge {
pub fn new() -> Self {
Self::default()
}
}
impl std::fmt::Display for ResponseBodyTooLarge {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "response body exceeded maximum allowed 4 MiB limit")
}
}
impl std::error::Error for ResponseBodyTooLarge {}
#[cfg(feature = "reqwest")]
mod reqwest {
use opentelemetry::otel_debug;
use crate::ResponseBodyTooLarge;
use super::{
async_trait, Bytes, HttpClient, HttpError, Request, Response, MAX_RESPONSE_BODY_BYTES,
};
#[async_trait]
impl HttpClient for reqwest::Client {
async fn send_bytes(&self, request: Request<Bytes>) -> Result<Response<Bytes>, HttpError> {
otel_debug!(name: "ReqwestClient.Send");
let request = request.try_into()?;
let mut response = self.execute(request).await?;
let capacity = response
.content_length()
.unwrap_or(0)
.min(MAX_RESPONSE_BODY_BYTES as u64) as usize;
let mut body_bytes = bytes::BytesMut::with_capacity(capacity);
let status = response.status();
let headers = std::mem::take(response.headers_mut());
while let Some(chunk) = response.chunk().await? {
if body_bytes.len() + chunk.len() > MAX_RESPONSE_BODY_BYTES {
return Err(Box::new(ResponseBodyTooLarge::new()));
}
body_bytes.extend_from_slice(&chunk);
}
let mut http_response = Response::builder()
.status(status)
.body(body_bytes.freeze())?;
*http_response.headers_mut() = headers;
Ok(http_response)
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg(feature = "reqwest-blocking")]
#[async_trait]
impl HttpClient for reqwest::blocking::Client {
async fn send_bytes(&self, request: Request<Bytes>) -> Result<Response<Bytes>, HttpError> {
use std::io::Read;
otel_debug!(name: "ReqwestBlockingClient.Send");
let request = request.try_into()?;
let mut response = self.execute(request)?;
let capacity = response
.content_length()
.unwrap_or(0)
.min(MAX_RESPONSE_BODY_BYTES as u64) as usize;
let status = response.status();
let headers = std::mem::take(response.headers_mut());
let mut body_bytes = Vec::with_capacity(capacity);
response
.take(MAX_RESPONSE_BODY_BYTES as u64 + 1)
.read_to_end(&mut body_bytes)?;
if body_bytes.len() > MAX_RESPONSE_BODY_BYTES {
return Err(Box::new(ResponseBodyTooLarge::new()));
}
let mut http_response = Response::builder()
.status(status)
.body(Bytes::from(body_bytes))?;
*http_response.headers_mut() = headers;
Ok(http_response)
}
}
}
#[cfg(feature = "hyper")]
pub mod hyper {
use super::{
async_trait, Bytes, HttpClient, HttpError, Request, Response, MAX_RESPONSE_BODY_BYTES,
};
use crate::ResponseBodyTooLarge;
use http::HeaderValue;
use http_body_util::{BodyExt, Full};
use hyper::body::Body as _;
use hyper_util::client::legacy::{
connect::{Connect, HttpConnector},
Client,
};
use opentelemetry::otel_debug;
use std::fmt::Debug;
use std::time::Duration;
use tokio::time;
#[derive(Debug, Clone)]
pub struct HyperClient<C = HttpConnector>
where
C: Connect + Clone + Send + Sync + 'static,
{
inner: Client<C, Full<Bytes>>,
timeout: Duration,
authorization: Option<HeaderValue>,
}
impl<C> HyperClient<C>
where
C: Connect + Clone + Send + Sync + 'static,
{
pub fn new(connector: C, timeout: Duration, authorization: Option<HeaderValue>) -> Self {
let inner = Client::builder(hyper_util::rt::TokioExecutor::new()).build(connector);
Self {
inner,
timeout,
authorization,
}
}
}
impl HyperClient<HttpConnector> {
pub fn with_default_connector(
timeout: Duration,
authorization: Option<HeaderValue>,
) -> Self {
Self::new(HttpConnector::new(), timeout, authorization)
}
}
#[async_trait]
impl<C> HttpClient for HyperClient<C>
where
C: Connect + Clone + Send + Sync + 'static,
HyperClient<C>: Debug,
{
async fn send_bytes(&self, request: Request<Bytes>) -> Result<Response<Bytes>, HttpError> {
otel_debug!(name: "HyperClient.Send");
let (parts, body) = request.into_parts();
let mut request = Request::from_parts(parts, Full::from(body));
if let Some(ref authorization) = self.authorization {
request
.headers_mut()
.insert(http::header::AUTHORIZATION, authorization.clone());
}
time::timeout(self.timeout, async {
let mut response = self.inner.request(request).await?;
let capacity = response
.body()
.size_hint()
.upper()
.unwrap_or(0)
.min(MAX_RESPONSE_BODY_BYTES as u64) as usize;
let mut body_bytes = bytes::BytesMut::with_capacity(capacity);
let status = response.status();
let headers = std::mem::take(response.headers_mut());
let mut body = response.into_body();
while let Some(frame) = body.frame().await {
let frame = frame?;
if let Ok(chunk) = frame.into_data() {
if body_bytes.len() + chunk.len() > MAX_RESPONSE_BODY_BYTES {
return Err(Box::new(ResponseBodyTooLarge::new()) as HttpError);
}
body_bytes.extend_from_slice(&chunk);
}
}
let mut http_response = Response::builder()
.status(status)
.body(body_bytes.freeze())?;
*http_response.headers_mut() = headers;
Ok(http_response)
})
.await?
}
}
}
mod private {
pub trait Sealed {}
impl<T> Sealed for http::Response<T> {}
}
pub trait ResponseExt: private::Sealed + Sized {
fn error_for_status(self) -> Result<Self, HttpError>;
}
impl<T> ResponseExt for Response<T> {
fn error_for_status(self) -> Result<Self, HttpError> {
if self.status().is_success() {
Ok(self)
} else {
Err(format!("request failed with status {}", self.status()).into())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use http::HeaderValue;
#[test]
fn response_body_too_large_construction() {
for error in [ResponseBodyTooLarge::new(), ResponseBodyTooLarge::default()] {
let error: HttpError = Box::new(error);
assert!(error.downcast_ref::<ResponseBodyTooLarge>().is_some());
assert_eq!(
error.to_string(),
"response body exceeded maximum allowed 4 MiB limit"
);
}
}
#[cfg(all(
any(feature = "hyper", feature = "reqwest", feature = "reqwest-blocking"),
not(target_arch = "wasm32")
))]
use std::io::{Read, Write};
#[cfg(all(
any(feature = "hyper", feature = "reqwest", feature = "reqwest-blocking"),
not(target_arch = "wasm32")
))]
use std::net::{SocketAddr, TcpListener};
#[cfg(all(
any(feature = "hyper", feature = "reqwest", feature = "reqwest-blocking"),
not(target_arch = "wasm32")
))]
use std::thread::JoinHandle;
#[test]
fn http_headers_get() {
let mut carrier = http::HeaderMap::new();
HeaderInjector(&mut carrier).set("headerName", "value".to_string());
assert_eq!(
HeaderExtractor(&carrier).get("HEADERNAME"),
Some("value"),
"case insensitive extraction"
)
}
#[test]
fn http_headers_get_all() {
let mut carrier = http::HeaderMap::new();
carrier.append("headerName", HeaderValue::from_static("value"));
carrier.append("headerName", HeaderValue::from_static("value2"));
carrier.append("headerName", HeaderValue::from_static("value3"));
assert_eq!(
HeaderExtractor(&carrier).get_all("HEADERNAME"),
Some(vec!["value", "value2", "value3"]),
"all values from a key extraction"
)
}
#[test]
fn http_headers_get_all_missing_key() {
let mut carrier = http::HeaderMap::new();
carrier.append("headerName", HeaderValue::from_static("value"));
assert_eq!(
HeaderExtractor(&carrier).get_all("not_existing"),
None,
"all values from a missing key extraction"
)
}
#[test]
fn http_headers_keys() {
let mut carrier = http::HeaderMap::new();
HeaderInjector(&mut carrier).set("headerName1", "value1".to_string());
HeaderInjector(&mut carrier).set("headerName2", "value2".to_string());
let extractor = HeaderExtractor(&carrier);
let got = extractor.keys();
assert_eq!(got.len(), 2);
assert!(got.contains(&"headername1"));
assert!(got.contains(&"headername2"));
}
#[test]
fn http_headers_reserve() {
let mut carrier = http::HeaderMap::new();
{
let mut injector = HeaderInjector(&mut carrier);
injector.reserve(10);
injector.set("test-header", "test-value".to_string());
}
assert_eq!(
HeaderExtractor(&carrier).get("test-header"),
Some("test-value")
);
{
let mut injector = HeaderInjector(&mut carrier);
injector.reserve(0);
injector.set("another-header", "another-value".to_string());
}
assert_eq!(
HeaderExtractor(&carrier).get("another-header"),
Some("another-value")
);
let mut new_carrier = http::HeaderMap::new();
{
let mut new_injector = HeaderInjector(&mut new_carrier);
new_injector.reserve(5);
}
let initial_capacity = new_carrier.capacity();
{
let mut new_injector = HeaderInjector(&mut new_carrier);
for i in 0..3 {
new_injector.set(&format!("header-{}", i), format!("value-{}", i));
}
}
assert!(new_carrier.capacity() >= initial_capacity);
assert!(new_carrier.capacity() >= 5);
}
#[test]
fn error_for_status_matches_http_status_class() {
for status in [http::StatusCode::OK, http::StatusCode::NO_CONTENT] {
let response = Response::builder().status(status).body(()).unwrap();
assert!(response.error_for_status().is_ok());
}
for status in [
http::StatusCode::MOVED_PERMANENTLY,
http::StatusCode::BAD_REQUEST,
http::StatusCode::TOO_MANY_REQUESTS,
http::StatusCode::INTERNAL_SERVER_ERROR,
] {
let response = Response::builder().status(status).body(()).unwrap();
assert!(response.error_for_status().is_err());
}
}
#[cfg(all(
any(feature = "hyper", feature = "reqwest", feature = "reqwest-blocking"),
not(target_arch = "wasm32")
))]
fn spawn_error_response_server() -> (SocketAddr, JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut request = [0; 1024];
let _ = stream.read(&mut request).unwrap();
stream
.write_all(
b"HTTP/1.1 429 Too Many Requests\r\n\
Retry-After: 7\r\n\
Content-Length: 0\r\n\
Connection: close\r\n\r\n",
)
.unwrap();
});
(address, server)
}
#[cfg(all(feature = "reqwest-blocking", not(target_arch = "wasm32")))]
#[test]
fn reqwest_blocking_preserves_error_response_status_and_headers() {
let (address, server) = spawn_error_response_server();
let client = ::reqwest::blocking::Client::new();
let request = Request::post(format!("http://{address}/v1/traces"))
.body(Bytes::new())
.unwrap();
let response = futures_executor::block_on(client.send_bytes(request)).unwrap();
server.join().unwrap();
assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS);
assert_eq!(response.headers().get("retry-after").unwrap(), "7");
}
#[cfg(all(feature = "reqwest", not(target_arch = "wasm32")))]
#[test]
fn reqwest_async_preserves_error_response_status_and_headers() {
let (address, server) = spawn_error_response_server();
let client = ::reqwest::Client::new();
let request = Request::post(format!("http://{address}/v1/traces"))
.body(Bytes::new())
.unwrap();
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let response = runtime.block_on(client.send_bytes(request)).unwrap();
server.join().unwrap();
assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS);
assert_eq!(response.headers().get("retry-after").unwrap(), "7");
}
#[cfg(all(feature = "hyper", not(target_arch = "wasm32")))]
#[test]
fn hyper_preserves_error_response_status_and_headers() {
let (address, server) = spawn_error_response_server();
let client = crate::hyper::HyperClient::with_default_connector(
std::time::Duration::from_secs(2),
None,
);
let request = Request::post(format!("http://{address}/v1/traces"))
.body(Bytes::new())
.unwrap();
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let response = runtime.block_on(client.send_bytes(request)).unwrap();
server.join().unwrap();
assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS);
assert_eq!(response.headers().get("retry-after").unwrap(), "7");
}
#[cfg(all(
test,
any(feature = "reqwest", feature = "reqwest-blocking", feature = "hyper")
))]
mod body_limit_tests {
use super::MAX_RESPONSE_BODY_BYTES;
use crate::HttpClient;
use bytes::Bytes;
use http::Request;
#[cfg(feature = "hyper")]
use std::future::Future;
use std::net::SocketAddr;
#[cfg(feature = "hyper")]
use std::pin::Pin;
#[cfg(feature = "hyper")]
use std::task::{Context, Poll};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[cfg(feature = "hyper")]
#[derive(Clone, Debug)]
struct LocalConnector(SocketAddr);
#[cfg(feature = "hyper")]
impl tower_service::Service<http::Uri> for LocalConnector {
type Response = hyper_util::rt::TokioIo<tokio::net::TcpStream>;
type Error = std::io::Error;
type Future =
Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
fn poll_ready(&mut self, _context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _uri: http::Uri) -> Self::Future {
let address = self.0;
Box::pin(async move {
tokio::net::TcpStream::connect(address)
.await
.map(hyper_util::rt::TokioIo::new)
})
}
}
async fn start_server(body_size: usize) -> SocketAddr {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let body = vec![b'a'; body_size];
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.write_all(&body).await;
let _ = socket.shutdown().await;
}
});
addr
}
#[cfg(feature = "hyper")]
async fn start_stalled_body_server() -> SocketAddr {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let _ = socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 1\r\nConnection: close\r\n\r\n",
)
.await;
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
}
});
addr
}
async fn assert_body_size(client: &dyn HttpClient, addr: SocketAddr, expected_size: usize) {
let request = Request::builder()
.method("POST")
.uri(format!("http://{}/", addr))
.body(Bytes::new())
.unwrap();
let response = client.send_bytes(request).await.unwrap();
assert_eq!(response.body().len(), expected_size);
}
async fn assert_exceeds_limit(client: &dyn HttpClient, addr: SocketAddr) {
let request = Request::builder()
.method("POST")
.uri(format!("http://{}/", addr))
.body(Bytes::new())
.unwrap();
let error = client.send_bytes(request).await.unwrap_err();
assert!(error
.downcast_ref::<crate::ResponseBodyTooLarge>()
.is_some());
}
#[cfg(feature = "reqwest-blocking")]
fn start_blocking_server(body_size: usize) -> SocketAddr {
use std::io::{Read, Write};
use std::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
if let Ok((mut socket, _)) = listener.accept() {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf);
let body = vec![b'a'; body_size];
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = socket.write_all(response.as_bytes());
let _ = socket.write_all(&body);
}
});
addr
}
#[cfg(feature = "reqwest")]
#[tokio::test]
async fn reqwest_body_within_limit() {
let addr = start_server(MAX_RESPONSE_BODY_BYTES).await;
assert_body_size(&reqwest::Client::new(), addr, MAX_RESPONSE_BODY_BYTES).await;
}
#[cfg(feature = "reqwest")]
#[tokio::test]
async fn reqwest_body_exceeds_limit() {
let addr = start_server(MAX_RESPONSE_BODY_BYTES + 1).await;
assert_exceeds_limit(&reqwest::Client::new(), addr).await;
}
#[cfg(feature = "reqwest-blocking")]
#[test]
fn reqwest_blocking_body_within_limit() {
let addr = start_blocking_server(MAX_RESPONSE_BODY_BYTES);
futures_executor::block_on(assert_body_size(
&reqwest::blocking::Client::new(),
addr,
MAX_RESPONSE_BODY_BYTES,
));
}
#[cfg(feature = "reqwest-blocking")]
#[test]
fn reqwest_blocking_body_exceeds_limit() {
let addr = start_blocking_server(MAX_RESPONSE_BODY_BYTES + 1);
futures_executor::block_on(assert_exceeds_limit(
&reqwest::blocking::Client::new(),
addr,
));
}
#[cfg(feature = "hyper")]
#[tokio::test]
async fn hyper_body_within_limit() {
let addr = start_server(MAX_RESPONSE_BODY_BYTES).await;
let client = crate::hyper::HyperClient::with_default_connector(
std::time::Duration::from_secs(5),
None,
);
assert_body_size(&client, addr, MAX_RESPONSE_BODY_BYTES).await;
}
#[cfg(feature = "hyper")]
#[tokio::test]
async fn hyper_client_new_accepts_custom_connector() {
let addr = start_server(100).await;
let client = crate::hyper::HyperClient::new(
LocalConnector(addr),
std::time::Duration::from_secs(5),
None,
);
assert_body_size(&client, addr, 100).await;
}
#[cfg(feature = "hyper")]
#[tokio::test]
async fn hyper_body_exceeds_limit() {
let addr = start_server(MAX_RESPONSE_BODY_BYTES + 1).await;
let client = crate::hyper::HyperClient::with_default_connector(
std::time::Duration::from_secs(5),
None,
);
assert_exceeds_limit(&client, addr).await;
}
#[cfg(feature = "hyper")]
#[tokio::test]
async fn hyper_timeout_covers_response_body() {
let addr = start_stalled_body_server().await;
let client = crate::hyper::HyperClient::with_default_connector(
std::time::Duration::from_millis(25),
None,
);
let request = Request::post(format!("http://{addr}/"))
.body(Bytes::new())
.unwrap();
let error = tokio::time::timeout(
std::time::Duration::from_millis(200),
client.send_bytes(request),
)
.await
.expect("HyperClient must enforce its configured timeout")
.expect_err("stalled response body must time out");
assert!(error
.downcast_ref::<tokio::time::error::Elapsed>()
.is_some());
}
}
}