Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
312 changes: 176 additions & 136 deletions rust/networking/src/reqwest.rs
Original file line number Diff line number Diff line change
@@ -1,136 +1,176 @@
//! An [`http::Client`] implementation that utilizes [`reqwest`].

use ::http::{HeaderName, HeaderValue};
use async_trait::async_trait;
use reqwest::{Certificate, RequestBuilder};
use std::collections::HashMap;
use std::str::FromStr;
use std::time::Duration;
use tracing::warn;

use crate::http;

/// Options for configuring the [`reqwest`] [`Client`].
#[derive(Debug, Clone)]
pub struct ClientOptions<'a> {
pub additional_root_certs: Vec<Certificate>,
pub timeout: Duration,
pub default_headers: HashMap<&'a str, &'a str>,
}

impl<'a> Default for ClientOptions<'a> {
fn default() -> Self {
Self {
additional_root_certs: Vec::new(),
timeout: Duration::from_secs(30),
default_headers: HashMap::from([(
"User-Agent",
concat!(env!("CARGO_PKG_NAME"), "/", env!("CARGO_PKG_VERSION")),
)]),
}
}
}

/// An [`http::Client`] implementation that utilizes [`reqwest`].
#[derive(Clone, Debug, Default)]
pub struct Client {
// reqwest::Client holds a connection pool. It's reference-counted
// internally, so this field is relatively cheap to clone.
http: reqwest::Client,
}

impl Client {
pub fn new(options: ClientOptions) -> Self {
let mut b = reqwest::Client::builder()
.timeout(options.timeout)
// The service checker needs access to the server's certificate to
// warn if it will expire soon.
.tls_info(true)
.use_rustls_tls();

let mut default_headers = reqwest::header::HeaderMap::new();
for (key, value) in options.default_headers {
if let (Ok(header_name), Ok(header_value)) =
(HeaderName::from_str(key), HeaderValue::from_str(value))
{
default_headers.append(header_name, header_value);
}
}
b = b.default_headers(default_headers);

for c in options.additional_root_certs {
b = b.add_root_certificate(c);
}
Self {
http: b.build().expect("TODO"),
}
}

pub fn to_reqwest(&self, request: http::Request) -> RequestBuilder {
let mut request_builder = match request.method {
http::Method::Get => self.http.get(request.url),
http::Method::Put => self.http.put(request.url),
http::Method::Post => self.http.post(request.url),
http::Method::Delete => self.http.delete(request.url),
};

let mut headers = reqwest::header::HeaderMap::new();
for (key, value) in request.headers {
if let (Ok(header_name), Ok(header_value)) =
(HeaderName::from_str(&key), HeaderValue::from_str(&value))
{
headers.append(header_name, header_value);
}
}
request_builder = request_builder.headers(headers);

if let Some(body) = request.body {
request_builder = request_builder.body(body);
}

if let Some(timeout) = request.timeout {
request_builder = request_builder.timeout(timeout);
}
request_builder
}

pub async fn to_response(
&self,
resp: Result<reqwest::Response, reqwest::Error>,
) -> Result<http::Response, reqwest::Error> {
match resp {
Err(err) => {
warn!(%err, "error sending HTTP request");
Err(err)
}
Ok(response) => {
let status = response.status().as_u16();
let mut headers = HashMap::new();
for (header_name, header_value) in response.headers() {
if let Ok(value) = header_value.to_str() {
headers.insert(header_name.to_string(), value.to_owned());
}
}
match response.bytes().await {
Err(err) => {
warn!(%err, "error receiving HTTP response");
Err(err)
}
Ok(bytes) => Ok(http::Response {
status_code: status,
headers,
body: bytes.to_vec(),
}),
}
}
}
}
}

#[async_trait]
impl http::Client for Client {
async fn send(&self, request: http::Request) -> Option<http::Response> {
let resp = self.to_reqwest(request).send().await;
self.to_response(resp).await.ok()
}
}
//! An [`http::Client`] implementation that utilizes [`reqwest`].

use ::http::{HeaderName, HeaderValue};
use async_trait::async_trait;
use reqwest::{Certificate, RequestBuilder};
use std::collections::HashMap;
use std::str::FromStr;
use std::time::Duration;
use tracing::warn;

use crate::http;

/// Options for configuring the [`reqwest`] [`Client`].
#[derive(Debug, Clone)]
pub struct ClientOptions<'a> {
pub additional_root_certs: Vec<Certificate>,
pub timeout: Duration,
pub default_headers: HashMap<&'a str, &'a str>,
}

/// Upper bound on the accepted HTTP response body size, in bytes.
///
/// A misbehaving or compromised realm can stream an unbounded body and
/// otherwise exhaust the client's memory; the request timeout bounds duration,
/// not size. This cap is enforced while streaming, before the full body is
/// buffered.
const MAX_RESPONSE_BODY_BYTES: usize = 16 * 1024 * 1024;

/// Error variants produced while accumulating a capped response body.
#[derive(Debug, thiserror::Error)]
enum ResponseBodyError {
#[error("HTTP response body too large: {0} bytes (limit {MAX_RESPONSE_BODY_BYTES})")]
TooLarge(usize),
#[error("error receiving HTTP response body: {0}")]
Receive(#[from] reqwest::Error),
}

/// Streams the response body into memory, aborting once the size cap is hit.
async fn read_body_capped(mut response: reqwest::Response) -> Result<Vec<u8>, ResponseBodyError> {
if let Some(length) = response.content_length() {
if length as usize > MAX_RESPONSE_BODY_BYTES {
return Err(ResponseBodyError::TooLarge(length as usize));
}
}

let capacity = usize::try_from(response.content_length().unwrap_or(0))
.unwrap_or(0)
.min(MAX_RESPONSE_BODY_BYTES);
let mut body = Vec::with_capacity(capacity);

while let Some(chunk) = response.chunk().await? {
body.extend_from_slice(&chunk);
if body.len() > MAX_RESPONSE_BODY_BYTES {
return Err(ResponseBodyError::TooLarge(body.len()));
}
}

Ok(body)
}

impl<'a> Default for ClientOptions<'a> {
fn default() -> Self {
Self {
additional_root_certs: Vec::new(),
timeout: Duration::from_secs(30),
default_headers: HashMap::from([(
"User-Agent",
concat!(env!("CARGO_PKG_NAME"), "/", env!("CARGO_PKG_VERSION")),
)]),
}
}
}

/// An [`http::Client`] implementation that utilizes [`reqwest`].
#[derive(Clone, Debug, Default)]
pub struct Client {
// reqwest::Client holds a connection pool. It's reference-counted
// internally, so this field is relatively cheap to clone.
http: reqwest::Client,
}

impl Client {
pub fn new(options: ClientOptions) -> Self {
let mut b = reqwest::Client::builder()
.timeout(options.timeout)
// The service checker needs access to the server's certificate to
// warn if it will expire soon.
.tls_info(true)
.use_rustls_tls();

let mut default_headers = reqwest::header::HeaderMap::new();
for (key, value) in options.default_headers {
if let (Ok(header_name), Ok(header_value)) =
(HeaderName::from_str(key), HeaderValue::from_str(value))
{
default_headers.append(header_name, header_value);
}
}
b = b.default_headers(default_headers);

for c in options.additional_root_certs {
b = b.add_root_certificate(c);
}
Self {
http: b.build().expect("TODO"),
}
}

pub fn to_reqwest(&self, request: http::Request) -> RequestBuilder {
let mut request_builder = match request.method {
http::Method::Get => self.http.get(request.url),
http::Method::Put => self.http.put(request.url),
http::Method::Post => self.http.post(request.url),
http::Method::Delete => self.http.delete(request.url),
};

let mut headers = reqwest::header::HeaderMap::new();
for (key, value) in request.headers {
if let (Ok(header_name), Ok(header_value)) =
(HeaderName::from_str(&key), HeaderValue::from_str(&value))
{
headers.append(header_name, header_value);
}
}
request_builder = request_builder.headers(headers);

if let Some(body) = request.body {
request_builder = request_builder.body(body);
}

if let Some(timeout) = request.timeout {
request_builder = request_builder.timeout(timeout);
}
request_builder
}

pub async fn to_response(
&self,
resp: Result<reqwest::Response, reqwest::Error>,
) -> Option<http::Response> {
match resp {
Err(err) => {
warn!(%err, "error sending HTTP request");
None
}
Ok(response) => {
let status = response.status().as_u16();
let mut headers = HashMap::new();
for (header_name, header_value) in response.headers() {
if let Ok(value) = header_value.to_str() {
headers.insert(header_name.to_string(), value.to_owned());
}
}
match read_body_capped(response).await {
Err(err) => {
warn!(%err, "error receiving HTTP response");
None
}
Ok(bytes) => Some(http::Response {
status_code: status,
headers,
body: bytes,
}),
}
}
}
}
}

#[async_trait]
impl http::Client for Client {
async fn send(&self, request: http::Request) -> Option<http::Response> {
let resp = self.to_reqwest(request).send().await;
self.to_response(resp).await
}
}
Loading