use std::time::Duration;
use async_trait::async_trait;
use futures_util::TryStreamExt;
use reqwest::header::{self, HeaderValue};
use reqwest::{Client, ClientBuilder, Response, redirect};
use url::Url;
use crate::upstream::origins::{OriginKind, OriginSet};
use crate::upstream::resolver::GuardedResolver;
use crate::upstream::{
ArtifactBody, ArtifactRequest, ByteStream, MetadataRequest, MetadataResponse, Transport,
UpstreamError, UpstreamValidators, capped, collect_capped,
};
const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const METADATA_TOTAL_TIMEOUT: Duration = Duration::from_secs(30);
const ARTIFACT_READ_TIMEOUT: Duration = Duration::from_secs(30);
const ARTIFACT_TOTAL_TIMEOUT: Duration = Duration::from_secs(15 * 60);
const MAX_REDIRECT_HOPS: usize = 5;
const USER_AGENT: &str = concat!("probation/", env!("CARGO_PKG_VERSION"));
pub struct ReqwestTransport {
origins: OriginSet,
metadata_client: Client,
artifact_client: Client,
}
impl ReqwestTransport {
pub fn production() -> ReqwestTransport {
ReqwestTransport::for_origins(OriginSet::production())
}
pub fn for_origins(origins: OriginSet) -> ReqwestTransport {
let metadata_client = base_builder(&origins)
.timeout(METADATA_TOTAL_TIMEOUT)
.build()
.expect("the metadata client builds");
let artifact_client = base_builder(&origins)
.read_timeout(ARTIFACT_READ_TIMEOUT)
.timeout(ARTIFACT_TOTAL_TIMEOUT)
.no_gzip()
.no_brotli()
.no_zstd()
.build()
.expect("the artifact client builds");
ReqwestTransport {
origins,
metadata_client,
artifact_client,
}
}
fn admitted(&self, url: &Url) -> Result<OriginKind, UpstreamError> {
self.origins
.admit_any(url)
.map_err(UpstreamError::RejectedUrl)
}
}
fn base_builder(origins: &OriginSet) -> ClientBuilder {
Client::builder()
.tls_backend_rustls()
.user_agent(USER_AGENT)
.referer(false)
.no_proxy()
.connect_timeout(CONNECT_TIMEOUT)
.dns_resolver(GuardedResolver::new(origins.allows_private_addresses()))
.redirect(same_origin_policy(origins.clone()))
}
fn same_origin_policy(origins: OriginSet) -> redirect::Policy {
redirect::Policy::custom(move |attempt| {
if attempt.previous().len() > MAX_REDIRECT_HOPS {
return attempt.stop();
}
let Some(previous) = attempt.previous().last() else {
return attempt.stop();
};
let Some(kind) = origins.kind_of(previous) else {
return attempt.stop();
};
match origins.admit(attempt.url(), kind) {
Ok(()) => attempt.follow(),
Err(_) => attempt.stop(),
}
})
}
#[async_trait]
impl Transport for ReqwestTransport {
async fn fetch_metadata(
&self,
req: MetadataRequest,
) -> Result<MetadataResponse, UpstreamError> {
self.admitted(&req.url)?;
let mut request = self
.metadata_client
.get(req.url.clone())
.header(header::ACCEPT, req.accept);
if let Some(validators) = &req.validators {
if let Some(etag) = validator_header(validators.etag.as_deref()) {
request = request.header(header::IF_NONE_MATCH, etag);
}
if let Some(last_modified) = validator_header(validators.last_modified.as_deref()) {
request = request.header(header::IF_MODIFIED_SINCE, last_modified);
}
}
let response = request.send().await.map_err(upstream_error)?;
match response.status().as_u16() {
200 => {
let validators = validators_of(&response);
let body = collect_capped(byte_stream(response), req.max_bytes).await?;
Ok(MetadataResponse::Fresh { body, validators })
}
304 => Ok(MetadataResponse::NotModified {
validators: validators_of(&response),
}),
404 | 410 => Ok(MetadataResponse::Missing),
code => match refused_redirect(&response) {
Some(rejected) => Err(rejected),
None => Err(UpstreamError::Status(code)),
},
}
}
async fn open_artifact(&self, req: ArtifactRequest) -> Result<ArtifactBody, UpstreamError> {
self.admitted(&req.url)?;
let response = self
.artifact_client
.get(req.url.clone())
.send()
.await
.map_err(upstream_error)?;
let status = response.status().as_u16();
if status != 200 {
return match refused_redirect(&response) {
Some(rejected) => Err(rejected),
None => Err(UpstreamError::Status(status)),
};
}
let declared_length = response.content_length();
Ok(ArtifactBody {
declared_length,
stream: capped(byte_stream(response), req.max_bytes),
})
}
}
fn refused_redirect(response: &Response) -> Option<UpstreamError> {
if !matches!(response.status().as_u16(), 301 | 302 | 303 | 307 | 308) {
return None;
}
let location = response.headers().get(header::LOCATION)?.to_str().ok()?;
let to = response.url().join(location).ok()?;
Some(UpstreamError::RejectedRedirect { to })
}
fn validators_of(response: &Response) -> UpstreamValidators {
UpstreamValidators {
etag: header_string(response, header::ETAG),
last_modified: header_string(response, header::LAST_MODIFIED),
}
}
fn header_string(response: &Response, name: header::HeaderName) -> Option<String> {
response
.headers()
.get(name)
.and_then(|value| value.to_str().ok())
.map(str::to_owned)
}
fn validator_header(value: Option<&str>) -> Option<HeaderValue> {
value.and_then(|value| HeaderValue::from_str(value).ok())
}
fn byte_stream(response: Response) -> ByteStream {
Box::pin(response.bytes_stream().map_err(upstream_error))
}
fn upstream_error(err: reqwest::Error) -> UpstreamError {
if err.is_timeout() {
return UpstreamError::Timeout;
}
if let Some(found) = refusal_in(&err) {
return found;
}
UpstreamError::Transport(err.to_string())
}
pub fn refusal_in(err: &(dyn std::error::Error + 'static)) -> Option<UpstreamError> {
let mut current = Some(err);
while let Some(err) = current {
if let Some(upstream) = err.downcast_ref::<UpstreamError>() {
return Some(upstream.clone());
}
current = err.source();
}
None
}