use async_trait::async_trait;
use bytes::{Bytes, BytesMut};
use digest::{is_blob_url, is_manifest_digest_url};
use dragonfly_api::common::v2::{Download, Priority, SchedulingPolicy, TaskType};
use dragonfly_api::dfdaemon::v2::{
dfdaemon_upload_client::DfdaemonUploadClient as DfdaemonUploadGRPCClient, DownloadTaskRequest,
};
use dragonfly_api::scheduler::v2::scheduler_client::SchedulerClient;
use errors::{BackendError, DfdaemonError, Error, ProxyError};
use futures::{Stream, TryStreamExt};
use http::{default_proxy_rule_filtered_query_params, headermap_to_hashmap};
use id_generator::{IDGenerator, TaskIDParameter};
use net::{format_url, preferred_local_ip};
use pool::{Builder as PoolBuilder, Factory, Pool};
use reqwest::{
header::{HeaderMap, HeaderValue},
Client,
};
use reqwest_middleware::{ClientBuilder, ClientWithMiddleware};
use reqwest_tracing::TracingMiddleware;
use retry::{retry, retry_with_endpoints, RetryPolicy};
use rustls_pki_types::CertificateDer;
use selector::{SeedPeerSelector, Selector};
use std::collections::HashMap;
use std::net::IpAddr;
use std::str::FromStr;
use std::sync::Arc;
use std::time::Duration;
use tonic::transport::{Channel, Endpoint};
use tracing::debug;
#[cfg(feature = "preheat")]
use dragonfly_api::scheduler::v2::StatImageRequest as SchedulerStatImageRequest;
#[cfg(feature = "preheat")]
use oci_client::{
client::{current_platform_resolver, ClientConfig},
manifest::{
ImageIndexEntry, IMAGE_MANIFEST_LIST_MEDIA_TYPE, IMAGE_MANIFEST_MEDIA_TYPE,
OCI_IMAGE_INDEX_MEDIA_TYPE, OCI_IMAGE_MEDIA_TYPE,
},
secrets::RegistryAuth,
Client as OciClient, Reference, RegistryOperation,
};
#[cfg(feature = "preheat")]
use oci_spec::image::{Arch, Os};
#[cfg(feature = "preheat")]
use reqwest::header::{ACCEPT, AUTHORIZATION};
#[cfg(feature = "preheat")]
use tokio::sync::Semaphore;
use tokio::task::JoinSet;
use tracing::Instrument;
pub mod digest;
pub mod errors;
pub mod hashring;
pub mod id_generator;
pub use backon::ExponentialBuilder;
mod http;
mod net;
mod pool;
mod retry;
mod selector;
mod shutdown;
mod url;
const POOL_MAX_IDLE_PER_HOST: usize = 1024;
const POOL_IDLE_TIMEOUT: Duration = Duration::from_secs(90);
const CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
const KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(60);
const DEFAULT_CLIENT_POOL_IDLE_TIMEOUT: Duration = Duration::from_secs(30 * 60);
const DEFAULT_CLIENT_POOL_CAPACITY: usize = 128;
const DEFAULT_SCHEDULER_REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(10 * 60);
const DEFAULT_REPLICAS: usize = 2;
#[cfg(feature = "preheat")]
const STAT_IMAGE_SCOPE_ALL_SEED_PEERS: &str = "all_seed_peers";
#[cfg(feature = "preheat")]
const MIME_TYPES_DISTRIBUTION_MANIFEST: &[&str] = &[
IMAGE_MANIFEST_MEDIA_TYPE,
IMAGE_MANIFEST_LIST_MEDIA_TYPE,
OCI_IMAGE_MEDIA_TYPE,
OCI_IMAGE_INDEX_MEDIA_TYPE,
];
pub type Result<T> = std::result::Result<T, Error>;
pub type Body = Box<dyn Stream<Item = Result<Bytes>> + Send + Unpin>;
#[cfg(feature = "preheat")]
type PlatformResolver = Box<dyn Fn(&[ImageIndexEntry]) -> Option<String> + Send + Sync>;
#[async_trait]
pub trait Request {
async fn get(&self, request: &GetRequest) -> Result<GetResponse<Body>>;
async fn get_into(&self, request: &GetRequest, buf: &mut BytesMut) -> Result<GetResponse>;
#[cfg(feature = "preheat")]
async fn preheat_image(&self, request: &PreheatImageRequest) -> Result<()>;
#[cfg(feature = "preheat")]
async fn stat_image(&self, request: &StatImageRequest) -> Result<StatImageResponse>;
async fn preheat(&self, request: &PreheatRequest) -> Result<()>;
async fn lookup_endpoints(&self, request: &GetRequest) -> Result<Vec<String>>;
}
#[async_trait]
pub trait RequestWithEndpoints {
async fn get(&self, request: &GetRequest) -> Result<GetResponse<Body>>;
async fn get_into(&self, request: &GetRequest, buf: &mut BytesMut) -> Result<GetResponse>;
}
pub struct GetRequest {
pub url: String,
pub header: HeaderMap,
pub piece_length: Option<u64>,
pub tag: Option<String>,
pub application: Option<String>,
pub filtered_query_params: Vec<String>,
pub content_for_calculating_task_id: Option<String>,
pub enable_task_id_based_blob_digest: bool,
pub priority: Option<i32>,
pub replicas: usize,
pub timeout: Duration,
pub client_cert: Option<Vec<CertificateDer<'static>>>,
}
impl Default for GetRequest {
fn default() -> Self {
Self {
url: String::new(),
header: HeaderMap::new(),
piece_length: None,
tag: None,
application: None,
filtered_query_params: default_proxy_rule_filtered_query_params(),
content_for_calculating_task_id: None,
enable_task_id_based_blob_digest: true,
priority: None,
replicas: DEFAULT_REPLICAS,
timeout: DEFAULT_REQUEST_TIMEOUT,
client_cert: None,
}
}
}
impl GetRequest {
fn validate(&self) -> Result<()> {
if self.replicas == 0 {
return Err(Error::InvalidArgument(
"replicas must be positive".to_string(),
));
}
Ok(())
}
}
pub struct GetResponse<R = Body> {
pub success: bool,
pub header: HeaderMap,
pub status_code: Option<reqwest::StatusCode>,
pub body: Option<R>,
}
#[cfg(feature = "preheat")]
pub struct PreheatImageRequest {
pub image: String,
pub username: Option<String>,
pub password: Option<String>,
pub platform: Option<String>,
pub piece_length: Option<u64>,
pub tag: Option<String>,
pub application: Option<String>,
pub filtered_query_params: Vec<String>,
pub content_for_calculating_task_id: Option<String>,
pub enable_task_id_based_blob_digest: bool,
pub priority: Option<i32>,
pub replicas: usize,
pub timeout: Duration,
pub concurrent_task_count: usize,
pub client_cert: Option<Vec<CertificateDer<'static>>>,
}
#[cfg(feature = "preheat")]
impl Default for PreheatImageRequest {
fn default() -> Self {
Self {
image: String::new(),
username: None,
password: None,
platform: None,
piece_length: None,
tag: None,
application: None,
filtered_query_params: default_proxy_rule_filtered_query_params(),
content_for_calculating_task_id: None,
enable_task_id_based_blob_digest: true,
priority: None,
replicas: DEFAULT_REPLICAS,
timeout: DEFAULT_REQUEST_TIMEOUT,
concurrent_task_count: 4,
client_cert: None,
}
}
}
#[cfg(feature = "preheat")]
impl PreheatImageRequest {
fn validate(&self) -> Result<()> {
if self.replicas == 0 {
return Err(Error::InvalidArgument(
"replicas must be positive".to_string(),
));
}
if self.concurrent_task_count == 0 {
return Err(Error::InvalidArgument(
"concurrent task count must be positive".to_string(),
));
}
Ok(())
}
}
pub struct PreheatRequest {
pub url: String,
pub header: HeaderMap,
pub piece_length: Option<u64>,
pub tag: Option<String>,
pub application: Option<String>,
pub filtered_query_params: Vec<String>,
pub content_for_calculating_task_id: Option<String>,
pub enable_task_id_based_blob_digest: bool,
pub priority: Option<i32>,
pub replicas: usize,
pub timeout: Duration,
pub client_cert: Option<Vec<CertificateDer<'static>>>,
}
impl Default for PreheatRequest {
fn default() -> Self {
Self {
url: String::new(),
header: HeaderMap::new(),
piece_length: None,
tag: None,
application: None,
filtered_query_params: default_proxy_rule_filtered_query_params(),
content_for_calculating_task_id: None,
enable_task_id_based_blob_digest: true,
priority: None,
replicas: DEFAULT_REPLICAS,
timeout: DEFAULT_REQUEST_TIMEOUT,
client_cert: None,
}
}
}
impl PreheatRequest {
fn validate(&self) -> Result<()> {
if self.replicas == 0 {
return Err(Error::InvalidArgument(
"replicas must be positive".to_string(),
));
}
Ok(())
}
}
#[cfg(feature = "preheat")]
pub struct StatImageRequest {
pub image: String,
pub username: Option<String>,
pub password: Option<String>,
pub platform: Option<String>,
pub piece_length: Option<u64>,
pub tag: Option<String>,
pub application: Option<String>,
pub filtered_query_params: Vec<String>,
pub enable_task_id_based_blob_digest: bool,
pub timeout: Duration,
}
#[cfg(feature = "preheat")]
impl Default for StatImageRequest {
fn default() -> Self {
Self {
image: String::new(),
username: None,
password: None,
platform: None,
piece_length: None,
tag: None,
application: None,
filtered_query_params: default_proxy_rule_filtered_query_params(),
enable_task_id_based_blob_digest: true,
timeout: DEFAULT_REQUEST_TIMEOUT,
}
}
}
#[cfg(feature = "preheat")]
pub struct StatImageResponse {
pub layers: Vec<String>,
pub peers: Vec<PeerImage>,
}
#[cfg(feature = "preheat")]
pub struct PeerImage {
pub ip: String,
pub hostname: String,
pub cached_layers: Vec<Layer>,
}
#[cfg(feature = "preheat")]
pub struct Layer {
pub url: String,
pub is_finished: bool,
}
#[derive(Debug, Clone, Default)]
struct HTTPClientFactory {}
#[async_trait]
impl Factory<String, ClientWithMiddleware> for HTTPClientFactory {
type Error = Error;
async fn make_client(&self, proxy_addr: &String) -> Result<ClientWithMiddleware> {
let client = Client::builder()
.hickory_dns(true)
.danger_accept_invalid_certs(true)
.pool_max_idle_per_host(POOL_MAX_IDLE_PER_HOST)
.pool_idle_timeout(POOL_IDLE_TIMEOUT)
.tcp_keepalive(KEEP_ALIVE_INTERVAL)
.connect_timeout(CONNECT_TIMEOUT)
.proxy(reqwest::Proxy::all(proxy_addr).map_err(|err| {
Error::Internal(format!("failed to set proxy {proxy_addr}: {err}"))
})?)
.build()
.map_err(|err| Error::Internal(format!("failed to build reqwest client: {err}")))?;
Ok(ClientBuilder::new(client)
.with(TracingMiddleware::default())
.build())
}
}
pub struct ProxyBuilder {
scheduler_endpoint: String,
scheduler_request_timeout: Duration,
health_check_interval: Duration,
retry: RetryPolicy,
}
impl Default for ProxyBuilder {
fn default() -> Self {
Self {
scheduler_endpoint: "".to_string(),
scheduler_request_timeout: DEFAULT_SCHEDULER_REQUEST_TIMEOUT,
health_check_interval: Duration::from_secs(60),
retry: RetryPolicy::default(),
}
}
}
impl ProxyBuilder {
pub fn scheduler_endpoint(mut self, endpoint: String) -> Self {
self.scheduler_endpoint = endpoint;
self
}
pub fn scheduler_request_timeout(mut self, timeout: Duration) -> Self {
self.scheduler_request_timeout = timeout;
self
}
pub fn health_check_interval(mut self, interval: Duration) -> Self {
self.health_check_interval = interval;
self
}
pub fn max_retries(mut self, retries: u8) -> Self {
self.retry.max_retries = retries;
self
}
pub fn backoff(mut self, backoff: ExponentialBuilder) -> Self {
self.retry.backoff = Some(backoff);
self
}
pub async fn build(self) -> Result<Proxy> {
self.validate()?;
let scheduler_channel = Endpoint::from_shared(self.scheduler_endpoint.to_string())
.map_err(|err| Error::InvalidArgument(err.to_string()))?
.connect_timeout(self.scheduler_request_timeout)
.timeout(self.scheduler_request_timeout)
.connect()
.await
.map_err(|err| {
Error::Internal(format!(
"failed to connect to scheduler {}: {}",
self.scheduler_endpoint, err
))
})?;
let scheduler_client = SchedulerClient::new(scheduler_channel);
let seed_peer_selector = Arc::new(
SeedPeerSelector::new(scheduler_client, self.health_check_interval)
.await
.map_err(|err| {
Error::Internal(format!("failed to create seed peer selector: {err}"))
})?,
);
let seed_peer_selector_clone = seed_peer_selector.clone();
tokio::spawn(async move {
seed_peer_selector_clone.run().await;
});
let local_ip = preferred_local_ip()
.ok_or_else(|| {
Error::Internal("failed to detect a preferred local IP address".to_string())
})?
.to_string();
let hostname = hostname::get()
.map_err(|err| Error::Internal(format!("failed to get hostname: {err}")))?
.to_string_lossy()
.to_string();
let id_generator = IDGenerator::new(local_ip, hostname, true);
let proxy = Proxy {
#[cfg(feature = "preheat")]
scheduler_endpoint: self.scheduler_endpoint,
seed_peer_selector,
retry: self.retry,
client_pool: Arc::new(
PoolBuilder::new(HTTPClientFactory::default())
.capacity(DEFAULT_CLIENT_POOL_CAPACITY)
.idle_timeout(DEFAULT_CLIENT_POOL_IDLE_TIMEOUT)
.build(),
),
id_generator: Arc::new(id_generator),
};
Ok(proxy)
}
fn validate(&self) -> Result<()> {
if let Err(err) = ::url::Url::parse(&self.scheduler_endpoint) {
return Err(Error::InvalidArgument(err.to_string()));
};
if self.scheduler_request_timeout.as_millis() < 100 {
return Err(Error::InvalidArgument(
"scheduler request timeout must be at least 100 milliseconds".to_string(),
));
}
if self.health_check_interval.as_secs() < 1 || self.health_check_interval.as_secs() > 600 {
return Err(Error::InvalidArgument(
"health check interval must be between 1 and 600 seconds".to_string(),
));
}
if self.retry.max_retries > 10 {
return Err(Error::InvalidArgument(
"max retries must be between 0 and 10".to_string(),
));
}
Ok(())
}
}
#[derive(Clone)]
pub struct Proxy {
#[cfg(feature = "preheat")]
scheduler_endpoint: String,
seed_peer_selector: Arc<SeedPeerSelector>,
retry: RetryPolicy,
client_pool: Arc<Pool<String, String, ClientWithMiddleware, HTTPClientFactory>>,
id_generator: Arc<IDGenerator>,
}
impl Proxy {
pub fn builder() -> ProxyBuilder {
ProxyBuilder::default()
}
}
#[async_trait]
impl Request for Proxy {
async fn get(&self, request: &GetRequest) -> Result<GetResponse> {
request.validate()?;
let response = self.try_send(request).await?;
let header = response.headers().clone();
let status_code = response.status();
let body: Body = Box::new(
response
.bytes_stream()
.map_err(|err| Error::Internal(err.to_string())),
);
Ok(GetResponse {
success: status_code.is_success(),
header,
status_code: Some(status_code),
body: Some(body),
})
}
async fn get_into(&self, request: &GetRequest, buf: &mut BytesMut) -> Result<GetResponse> {
request.validate()?;
let endpoints = self.lookup_proxy_endpoints(request).await?;
let (body, result) = retry_with_endpoints(
self.retry,
&endpoints,
std::mem::take(buf),
|mut body, endpoint| async move {
let result = match self.send_to(endpoint, request).await {
Ok(mut response) => {
let status = response.status();
let header = response.headers().clone();
let len = body.len();
if let Some(content_length) = response.content_length() {
body.reserve(content_length as usize);
}
loop {
match response.chunk().await {
Ok(Some(chunk)) => body.extend_from_slice(&chunk),
Ok(None) => {
break Ok(GetResponse {
success: status.is_success(),
header,
status_code: Some(status),
body: None,
})
}
Err(err) => {
body.truncate(len);
if err.is_timeout() {
break Err(Error::RequestTimeout(err.to_string()));
}
break Err(Error::Internal(format!(
"failed to read response body: {err}"
)));
}
}
}
}
Err(err) => Err(err),
};
(body, result)
},
)
.await;
*buf = body;
result
}
#[cfg(feature = "preheat")]
async fn preheat_image(&self, request: &PreheatImageRequest) -> Result<()> {
request.validate()?;
let oci_client = Self::oci_client(request.platform.clone())?;
let reference: Reference = request
.image
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid image reference: {err}")))?;
let auth = match (&request.username, &request.password) {
(Some(username), Some(password)) => {
RegistryAuth::Basic(username.clone(), password.clone())
}
_ => RegistryAuth::Anonymous,
};
let (manifest, manifest_digest) =
oci_client
.pull_image_manifest(&reference, &auth)
.await
.map_err(|err| Error::Internal(format!("failed to pull image manifest: {err}")))?;
debug!(
"pulled manifest for image {} with digest {}, layers: {}",
request.image,
manifest_digest,
manifest.layers.len()
);
let token = oci_client
.auth(&reference, &auth, RegistryOperation::Pull)
.await
.map_err(|err| Error::Internal(format!("failed to authenticate with registry: {err}")))?
.ok_or_else(|| {
Error::Internal("registry did not return authentication token".to_string())
})?;
let mut header = HeaderMap::new();
header.insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|err| Error::Internal(format!("invalid auth token: {err}")))?,
);
let mut manifest_header = header.clone();
manifest_header.insert(
ACCEPT,
HeaderValue::from_str(&MIME_TYPES_DISTRIBUTION_MANIFEST.join(","))
.map_err(|err| Error::Internal(format!("invalid accept header: {err}")))?,
);
let registry = Self::resolve_registry(&reference);
let repository = reference.repository();
let mut targets = Vec::with_capacity(manifest.layers.len() + 2);
targets.push((
Self::build_manifest_url(registry, repository, &manifest_digest),
manifest_header.clone(),
));
for digest in std::iter::once(&manifest.config.digest)
.chain(manifest.layers.iter().map(|layer| &layer.digest))
{
targets.push((
Self::build_blob_url(registry, repository, digest),
header.clone(),
));
}
let semaphore = Arc::new(Semaphore::new(request.concurrent_task_count));
let mut join_set: JoinSet<Result<()>> = JoinSet::new();
for (url, header) in targets {
let semaphore = semaphore.clone();
let proxy = self.clone();
let preheat_request = PreheatRequest {
url,
header,
piece_length: request.piece_length,
tag: request.tag.clone(),
application: request.application.clone(),
filtered_query_params: request.filtered_query_params.clone(),
content_for_calculating_task_id: request.content_for_calculating_task_id.clone(),
enable_task_id_based_blob_digest: request.enable_task_id_based_blob_digest,
priority: request.priority,
replicas: request.replicas,
timeout: request.timeout,
client_cert: request.client_cert.clone(),
};
join_set.spawn(
async move {
let _permit = semaphore
.acquire()
.await
.map_err(|err| Error::Internal(err.to_string()))?;
proxy.preheat(&preheat_request).await?;
debug!("preheated: {}", preheat_request.url);
Ok(())
}
.in_current_span(),
);
}
while let Some(result) = join_set.join_next().await {
result.map_err(|err| Error::Internal(err.to_string()))??;
}
debug!("preheat completed for image: {}", request.image);
Ok(())
}
#[cfg(feature = "preheat")]
async fn stat_image(&self, request: &StatImageRequest) -> Result<StatImageResponse> {
let reference: Reference = request
.image
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid image reference: {err}")))?;
let registry = Self::resolve_registry(&reference);
let stat_image_request = SchedulerStatImageRequest {
url: Self::build_manifest_url(
registry,
reference.repository(),
reference
.digest()
.or_else(|| reference.tag())
.unwrap_or("latest"),
),
piece_length: request.piece_length,
tag: request.tag.clone(),
application: request.application.clone(),
filtered_query_params: request.filtered_query_params.clone(),
username: request.username.clone(),
password: request.password.clone(),
platform: request.platform.clone(),
timeout: Some(
prost_wkt_types::Duration::try_from(request.timeout).map_err(|err| {
Error::InvalidArgument(format!("invalid request timeout: {err}"))
})?,
),
scope: STAT_IMAGE_SCOPE_ALL_SEED_PEERS.to_string(),
enable_task_id_based_blob_digest: request.enable_task_id_based_blob_digest,
..Default::default()
};
let channel = Channel::from_shared(self.scheduler_endpoint.clone())
.map_err(|err| Error::InvalidArgument(err.to_string()))?
.connect_timeout(request.timeout)
.timeout(request.timeout)
.connect()
.await
.map_err(|err| {
Error::Internal(format!(
"failed to connect to scheduler {}: {}",
self.scheduler_endpoint, err
))
})?;
let response = SchedulerClient::new(channel)
.max_decoding_message_size(i32::MAX as usize)
.max_encoding_message_size(i32::MAX as usize)
.stat_image(stat_image_request)
.await
.map_err(|err| match err.code() {
tonic::Code::InvalidArgument => Error::InvalidArgument(format!(
"failed to stat image {}: {}",
request.image, err
)),
_ => Error::Internal(format!("failed to stat image {}: {}", request.image, err)),
})?
.into_inner();
Ok(StatImageResponse {
layers: response
.image
.map(|image| image.layers.into_iter().map(|layer| layer.url).collect())
.unwrap_or_default(),
peers: response
.peers
.into_iter()
.map(|peer| PeerImage {
ip: peer.ip,
hostname: peer.hostname,
cached_layers: peer
.cached_layers
.into_iter()
.map(|layer| Layer {
url: layer.url,
is_finished: layer.is_finished.unwrap_or_default(),
})
.collect(),
})
.collect(),
})
}
async fn preheat(&self, request: &PreheatRequest) -> Result<()> {
request.validate()?;
let task_id = self
.id_generator
.task_id(
if let Some(content) = request.content_for_calculating_task_id.clone() {
TaskIDParameter::Content(content)
} else if request.enable_task_id_based_blob_digest && is_blob_url(&request.url) {
TaskIDParameter::BlobDigestBased(request.url.clone())
} else if request.enable_task_id_based_blob_digest
&& is_manifest_digest_url(&request.url)
{
TaskIDParameter::ManifestDigestBased(request.url.clone())
} else {
TaskIDParameter::URLBased {
url: request.url.clone(),
piece_length: request.piece_length,
tag: request.tag.clone(),
application: request.application.clone(),
filtered_query_params: request.filtered_query_params.clone(),
revision: None,
}
},
)
.map_err(|err| Error::Internal(format!("failed to generate task id: {err}")))?;
let seed_peers = self
.seed_peer_selector
.select(task_id.clone(), request.replicas as u32)
.await
.map_err(|err| {
Error::Internal(format!("failed to select seed peers from scheduler: {err}"))
})?;
debug!("task {} selected seed peers: {:?}", task_id, seed_peers);
if seed_peers.len() < request.replicas {
return Err(Error::Internal(format!(
"insufficient seed peers for {} replicas, {} available",
request.replicas,
seed_peers.len()
)));
}
let download_task_request = DownloadTaskRequest {
download: Some(Download {
url: request.url.clone(),
r#type: TaskType::Standard as i32,
tag: request.tag.clone(),
application: request.application.clone(),
priority: request.priority.unwrap_or(Priority::Level6 as i32),
filtered_query_params: request.filtered_query_params.clone(),
request_header: headermap_to_hashmap(&request.header),
piece_length: request.piece_length,
timeout: Some(
prost_wkt_types::Duration::try_from(request.timeout).map_err(|err| {
Error::InvalidArgument(format!("invalid request timeout: {err}"))
})?,
),
content_for_calculating_task_id: request.content_for_calculating_task_id.clone(),
remote_ip: preferred_local_ip().map(|ip| ip.to_string()),
enable_task_id_based_blob_digest: request.enable_task_id_based_blob_digest,
scheduling_policy: SchedulingPolicy::Always as i32,
..Default::default()
}),
};
let mut join_set: JoinSet<Result<()>> = JoinSet::new();
for peer in seed_peers.iter() {
let addr = format_url(
"http",
IpAddr::from_str(&peer.ip).map_err(|err| Error::Internal(err.to_string()))?,
peer.port as u16,
);
let download_task_request = download_task_request.clone();
let timeout = request.timeout;
let retry_policy = self.retry;
join_set.spawn(
async move {
retry(retry_policy, || async {
let channel = Channel::from_shared(addr.clone())
.map_err(|err| Error::InvalidArgument(err.to_string()))?
.connect_timeout(timeout)
.timeout(timeout)
.connect()
.await
.map_err(|err| {
Error::Internal(format!(
"failed to connect to seed peer {addr}: {err}"
))
})?;
let mut client = DfdaemonUploadGRPCClient::new(channel)
.max_decoding_message_size(usize::MAX)
.max_encoding_message_size(usize::MAX);
let mut response = client
.download_task(download_task_request.clone())
.await
.map_err(Error::from_status)?
.into_inner();
while response
.message()
.await
.map_err(Error::from_status)?
.is_some()
{}
Ok(())
})
.await
}
.in_current_span(),
);
}
while let Some(result) = join_set.join_next().await {
result.map_err(|err| Error::Internal(err.to_string()))??;
}
Ok(())
}
async fn lookup_endpoints(&self, request: &GetRequest) -> Result<Vec<String>> {
request.validate()?;
let task_id = self
.id_generator
.task_id(
if let Some(content) = request.content_for_calculating_task_id.clone() {
TaskIDParameter::Content(content)
} else if request.enable_task_id_based_blob_digest && is_blob_url(&request.url) {
TaskIDParameter::BlobDigestBased(request.url.clone())
} else if request.enable_task_id_based_blob_digest
&& is_manifest_digest_url(&request.url)
{
TaskIDParameter::ManifestDigestBased(request.url.clone())
} else {
TaskIDParameter::URLBased {
url: request.url.clone(),
piece_length: request.piece_length,
tag: request.tag.clone(),
application: request.application.clone(),
filtered_query_params: request.filtered_query_params.clone(),
revision: None,
}
},
)
.map_err(|err| Error::Internal(format!("failed to generate task id: {err}")))?;
let seed_peers = self
.seed_peer_selector
.select(task_id.clone(), request.replicas as u32)
.await
.map_err(|err| {
Error::Internal(format!("failed to select seed peers from scheduler: {err}"))
})?;
debug!("task {} selected seed peers: {:?}", task_id, seed_peers);
let mut addrs = Vec::with_capacity(seed_peers.len());
for peer in seed_peers.iter() {
addrs.push(format_url(
"http",
IpAddr::from_str(&peer.ip).map_err(|err| Error::Internal(err.to_string()))?,
peer.port as u16,
));
}
Ok(addrs)
}
}
impl Proxy {
async fn lookup_proxy_endpoints(&self, request: &GetRequest) -> Result<Vec<String>> {
let task_id = self
.id_generator
.task_id(
if let Some(content) = request.content_for_calculating_task_id.clone() {
TaskIDParameter::Content(content)
} else if request.enable_task_id_based_blob_digest && is_blob_url(&request.url) {
TaskIDParameter::BlobDigestBased(request.url.clone())
} else if request.enable_task_id_based_blob_digest
&& is_manifest_digest_url(&request.url)
{
TaskIDParameter::ManifestDigestBased(request.url.clone())
} else {
TaskIDParameter::URLBased {
url: request.url.clone(),
piece_length: request.piece_length,
tag: request.tag.clone(),
application: request.application.clone(),
filtered_query_params: request.filtered_query_params.clone(),
revision: None,
}
},
)
.map_err(|err| Error::Internal(format!("failed to generate task id: {err}")))?;
let seed_peers = self
.seed_peer_selector
.select(task_id.clone(), request.replicas as u32)
.await
.map_err(|err| {
Error::Internal(format!("failed to select seed peers from scheduler: {err}"))
})?;
debug!("task {} selected seed peers: {:?}", task_id, seed_peers);
let mut endpoints = Vec::with_capacity(seed_peers.len());
for peer in seed_peers.iter() {
endpoints.push(format_url(
"http",
IpAddr::from_str(&peer.ip).map_err(|err| Error::Internal(err.to_string()))?,
peer.proxy_port as u16,
));
}
Ok(endpoints)
}
async fn try_send(&self, request: &GetRequest) -> Result<reqwest::Response> {
let endpoints = self.lookup_proxy_endpoints(request).await?;
let ((), result) =
retry_with_endpoints(self.retry, &endpoints, (), |(), endpoint| async move {
((), self.send_to(endpoint, request).await)
})
.await;
result
}
async fn send_to(&self, endpoint: String, request: &GetRequest) -> Result<reqwest::Response> {
let entry = self.client_pool.entry(&endpoint, &endpoint).await?;
self.send(&entry.client, request).await
}
async fn send(
&self,
client: &ClientWithMiddleware,
request: &GetRequest,
) -> Result<reqwest::Response> {
let headers = self.make_request_headers(request)?;
let response = client
.get(&request.url)
.headers(headers)
.timeout(request.timeout)
.send()
.await
.map_err(|err| match err {
reqwest_middleware::Error::Reqwest(err) if err.is_timeout() => {
Error::RequestTimeout(err.to_string())
}
err => Error::Internal(err.to_string()),
})?;
let status = response.status();
if status.is_success() {
return Ok(response);
}
let response_headers = response.headers().clone();
let header_map = headermap_to_hashmap(&response_headers);
let message = response.text().await.ok();
let error_type = response_headers
.get("X-Dragonfly-Error-Type")
.and_then(|v| v.to_str().ok());
match error_type {
Some("backend") => Err(Error::BackendError(BackendError {
message,
header: header_map,
status_code: Some(status),
})),
Some("proxy") => Err(Error::ProxyError(ProxyError {
message,
header: header_map,
status_code: Some(status),
})),
Some("dfdaemon") => Err(Error::DfdaemonError(DfdaemonError { message })),
Some(other) => Err(Error::ProxyError(ProxyError {
message: Some(format!("unknown error type from proxy: {other}")),
header: header_map,
status_code: Some(status),
})),
None => Err(Error::ProxyError(ProxyError {
message: Some(format!("unexpected status code from proxy: {status}")),
header: header_map,
status_code: Some(status),
})),
}
}
fn make_request_headers(&self, request: &GetRequest) -> Result<HeaderMap> {
let mut headers = request.header.clone();
if let Some(piece_length) = request.piece_length {
headers.insert(
"X-Dragonfly-Piece-Length",
piece_length.to_string().parse().map_err(|err| {
Error::InvalidArgument(format!("invalid piece length: {err}"))
})?,
);
}
if let Some(tag) = request.tag.clone() {
headers.insert(
"X-Dragonfly-Tag",
tag.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid tag: {err}")))?,
);
}
if let Some(application) = request.application.clone() {
headers.insert(
"X-Dragonfly-Application",
application
.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid application: {err}")))?,
);
}
if let Some(content_for_calculating_task_id) =
request.content_for_calculating_task_id.clone()
{
headers.insert(
"X-Dragonfly-Content-For-Calculating-Task-ID",
content_for_calculating_task_id
.to_string()
.parse()
.map_err(|err| {
Error::InvalidArgument(format!(
"invalid content for calculating task id: {err}"
))
})?,
);
}
headers.insert(
"X-Dragonfly-Enable-Task-ID-Based-Blob-Digest",
request
.enable_task_id_based_blob_digest
.to_string()
.parse()
.map_err(|err| {
Error::InvalidArgument(format!(
"invalid enable task id based blob digest: {err}"
))
})?,
);
if let Some(priority) = request.priority {
headers.insert(
"X-Dragonfly-Priority",
priority
.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid priority: {err}")))?,
);
}
if !request.filtered_query_params.is_empty() {
let value = request.filtered_query_params.join(",");
headers.insert(
"X-Dragonfly-Filtered-Query-Params",
value.parse().map_err(|err| {
Error::InvalidArgument(format!("invalid filtered query params: {err}"))
})?,
);
}
headers.insert("X-Dragonfly-Use-P2P", HeaderValue::from_static("true"));
Ok(headers)
}
}
impl Proxy {
#[cfg(feature = "preheat")]
fn build_blob_url(registry: &str, repository: &str, digest: &str) -> String {
format!("https://{registry}/v2/{repository}/blobs/{digest}")
}
#[cfg(feature = "preheat")]
fn build_manifest_url(registry: &str, repository: &str, reference: &str) -> String {
format!("https://{registry}/v2/{repository}/manifests/{reference}")
}
#[cfg(feature = "preheat")]
fn resolve_registry(reference: &Reference) -> &str {
match reference.registry() {
"docker.io" => "registry-1.docker.io",
registry => registry,
}
}
#[cfg(feature = "preheat")]
fn platform_resolver(platform: &str) -> Result<PlatformResolver> {
let (os, arch) = platform
.split_once('/')
.map(|(os, arch)| (Os::from(os), Arch::from(arch)))
.ok_or_else(|| {
Error::InvalidArgument(format!("invalid platform format '{platform}', expected 'os/arch' (e.g., 'linux/amd64')"))
})?;
Ok(Box::new(move |manifests: &[ImageIndexEntry]| {
manifests
.iter()
.find(|entry| {
entry
.platform
.as_ref()
.is_some_and(|platform| platform.os == os && platform.architecture == arch)
})
.map(|entry| entry.digest.clone())
}))
}
#[cfg(feature = "preheat")]
fn oci_client(platform: Option<String>) -> Result<OciClient> {
let oci_config = ClientConfig {
platform_resolver: match platform {
Some(platform) => Some(Self::platform_resolver(&platform)?),
None => Some(Box::new(current_platform_resolver)),
},
..ClientConfig::default()
};
Ok(OciClient::new(oci_config))
}
}
pub struct ProxyWithEndpointsBuilder {
endpoints: Vec<String>,
retry: RetryPolicy,
}
impl Default for ProxyWithEndpointsBuilder {
fn default() -> Self {
Self {
endpoints: Vec::new(),
retry: RetryPolicy::default(),
}
}
}
impl ProxyWithEndpointsBuilder {
pub fn endpoints(mut self, endpoints: Vec<String>) -> Self {
self.endpoints = endpoints;
self
}
pub fn max_retries(mut self, retries: u8) -> Self {
self.retry.max_retries = retries;
self
}
pub fn backoff(mut self, backoff: ExponentialBuilder) -> Self {
self.retry.backoff = Some(backoff);
self
}
pub async fn build(self) -> Result<ProxyWithEndpoints> {
self.validate()?;
let factory = HTTPClientFactory::default();
let mut clients = HashMap::with_capacity(self.endpoints.len());
for endpoint in self.endpoints.iter() {
if clients.contains_key(endpoint) {
continue;
}
clients.insert(endpoint.clone(), factory.make_client(endpoint).await?);
}
Ok(ProxyWithEndpoints {
endpoints: self.endpoints,
retry: self.retry,
clients,
})
}
fn validate(&self) -> Result<()> {
if self.endpoints.is_empty() {
return Err(Error::InvalidArgument(
"endpoints must not be empty".to_string(),
));
}
if self.retry.max_retries > 10 {
return Err(Error::InvalidArgument(
"max retries must be between 0 and 10".to_string(),
));
}
Ok(())
}
}
#[derive(Clone)]
pub struct ProxyWithEndpoints {
endpoints: Vec<String>,
retry: RetryPolicy,
clients: HashMap<String, ClientWithMiddleware>,
}
impl ProxyWithEndpoints {
pub fn builder() -> ProxyWithEndpointsBuilder {
ProxyWithEndpointsBuilder::default()
}
async fn try_send(&self, request: &GetRequest) -> Result<reqwest::Response> {
let ((), result) =
retry_with_endpoints(self.retry, &self.endpoints, (), |(), endpoint| async move {
((), self.send_to(&endpoint, request).await)
})
.await;
result
}
async fn send_to(&self, endpoint: &str, request: &GetRequest) -> Result<reqwest::Response> {
let client = self
.clients
.get(endpoint)
.ok_or_else(|| Error::Internal(format!("no client for endpoint {endpoint}")))?;
self.send(client, request).await
}
async fn send(
&self,
client: &ClientWithMiddleware,
request: &GetRequest,
) -> Result<reqwest::Response> {
let headers = self.make_request_headers(request)?;
let response = client
.get(&request.url)
.headers(headers)
.timeout(request.timeout)
.send()
.await
.map_err(|err| match err {
reqwest_middleware::Error::Reqwest(err) if err.is_timeout() => {
Error::RequestTimeout(err.to_string())
}
err => Error::Internal(err.to_string()),
})?;
let status = response.status();
if status.is_success() {
return Ok(response);
}
let response_headers = response.headers().clone();
let header_map = headermap_to_hashmap(&response_headers);
let message = response.text().await.ok();
let error_type = response_headers
.get("X-Dragonfly-Error-Type")
.and_then(|v| v.to_str().ok());
match error_type {
Some("backend") => Err(Error::BackendError(BackendError {
message,
header: header_map,
status_code: Some(status),
})),
Some("proxy") => Err(Error::ProxyError(ProxyError {
message,
header: header_map,
status_code: Some(status),
})),
Some("dfdaemon") => Err(Error::DfdaemonError(DfdaemonError { message })),
Some(other) => Err(Error::ProxyError(ProxyError {
message: Some(format!("unknown error type from proxy: {other}")),
header: header_map,
status_code: Some(status),
})),
None => Err(Error::ProxyError(ProxyError {
message: Some(format!("unexpected status code from proxy: {status}")),
header: header_map,
status_code: Some(status),
})),
}
}
fn make_request_headers(&self, request: &GetRequest) -> Result<HeaderMap> {
let mut headers = request.header.clone();
if let Some(piece_length) = request.piece_length {
headers.insert(
"X-Dragonfly-Piece-Length",
piece_length.to_string().parse().map_err(|err| {
Error::InvalidArgument(format!("invalid piece length: {err}"))
})?,
);
}
if let Some(tag) = request.tag.clone() {
headers.insert(
"X-Dragonfly-Tag",
tag.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid tag: {err}")))?,
);
}
if let Some(application) = request.application.clone() {
headers.insert(
"X-Dragonfly-Application",
application
.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid application: {err}")))?,
);
}
if let Some(content_for_calculating_task_id) =
request.content_for_calculating_task_id.clone()
{
headers.insert(
"X-Dragonfly-Content-For-Calculating-Task-ID",
content_for_calculating_task_id
.to_string()
.parse()
.map_err(|err| {
Error::InvalidArgument(format!(
"invalid content for calculating task id: {err}"
))
})?,
);
}
headers.insert(
"X-Dragonfly-Enable-Task-ID-Based-Blob-Digest",
request
.enable_task_id_based_blob_digest
.to_string()
.parse()
.map_err(|err| {
Error::InvalidArgument(format!(
"invalid enable task id based blob digest: {err}"
))
})?,
);
if let Some(priority) = request.priority {
headers.insert(
"X-Dragonfly-Priority",
priority
.to_string()
.parse()
.map_err(|err| Error::InvalidArgument(format!("invalid priority: {err}")))?,
);
}
if !request.filtered_query_params.is_empty() {
let value = request.filtered_query_params.join(",");
headers.insert(
"X-Dragonfly-Filtered-Query-Params",
value.parse().map_err(|err| {
Error::InvalidArgument(format!("invalid filtered query params: {err}"))
})?,
);
}
headers.insert("X-Dragonfly-Use-P2P", HeaderValue::from_static("true"));
Ok(headers)
}
}
#[async_trait]
impl RequestWithEndpoints for ProxyWithEndpoints {
async fn get(&self, request: &GetRequest) -> Result<GetResponse> {
request.validate()?;
let response = self.try_send(request).await?;
let header = response.headers().clone();
let status_code = response.status();
let body: Body = Box::new(
response
.bytes_stream()
.map_err(|err| Error::Internal(err.to_string())),
);
Ok(GetResponse {
success: status_code.is_success(),
header,
status_code: Some(status_code),
body: Some(body),
})
}
async fn get_into(&self, request: &GetRequest, buf: &mut BytesMut) -> Result<GetResponse> {
request.validate()?;
let (body, result) = retry_with_endpoints(
self.retry,
&self.endpoints,
std::mem::take(buf),
|mut body, endpoint| async move {
let result = match self.send_to(&endpoint, request).await {
Ok(mut response) => {
let status = response.status();
let header = response.headers().clone();
let len = body.len();
if let Some(content_length) = response.content_length() {
body.reserve(content_length as usize);
}
loop {
match response.chunk().await {
Ok(Some(chunk)) => body.extend_from_slice(&chunk),
Ok(None) => {
break Ok(GetResponse {
success: status.is_success(),
header,
status_code: Some(status),
body: None,
})
}
Err(err) => {
body.truncate(len);
if err.is_timeout() {
break Err(Error::RequestTimeout(err.to_string()));
}
break Err(Error::Internal(format!(
"failed to read response body: {err}"
)));
}
}
}
}
Err(err) => Err(err),
};
(body, result)
},
)
.await;
*buf = body;
result
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Request;
use dragonfly_api::common::v2::Host;
use dragonfly_api::dfdaemon::v2::DownloadTaskResponse;
use dragonfly_api::scheduler::v2::ListHostsResponse;
use mocktail::prelude::*;
use std::time::Duration;
use tonic_health::pb::health_check_response::ServingStatus;
use tonic_health::pb::HealthCheckResponse;
async fn setup_mock_scheduler(hosts: Vec<Host>) -> Result<mocktail::server::MockServer> {
setup_mock_scheduler_with_mocks(hosts, MockSet::new()).await
}
async fn setup_mock_scheduler_with_mocks(
hosts: Vec<Host>,
mut mocks: MockSet,
) -> Result<mocktail::server::MockServer> {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/ListHosts");
then.pb(ListHostsResponse { hosts });
});
let server = MockServer::new_grpc("scheduler.v2.Scheduler").with_mocks(mocks);
server.start().await.map_err(|err| {
Error::Internal(format!("failed to start mock scheduler server: {err}"))
})?;
Ok(server)
}
async fn setup_mock_seed_peer(mut mocks: MockSet) -> Result<mocktail::server::MockServer> {
mocks.mock(|when, then| {
when.path("/grpc.health.v1.Health/Check");
then.pb(HealthCheckResponse {
status: ServingStatus::Serving as i32,
});
});
let server = MockServer::new_grpc("dfdaemon.v2.DfdaemonUpload").with_mocks(mocks);
server.start().await.map_err(|err| {
Error::Internal(format!("failed to start mock seed peer server: {err}"))
})?;
Ok(server)
}
fn create_seed_peer_host(name: &str, port: u16, proxy_port: u16) -> Host {
Host {
id: name.to_string(),
r#type: 1,
hostname: name.to_string(),
ip: "127.0.0.1".to_string(),
port: port as i32,
proxy_port: proxy_port as i32,
name: name.to_string(),
..Default::default()
}
}
async fn setup_mock_seed_peer_proxy(mocks: MockSet) -> Result<mocktail::server::MockServer> {
let server = MockServer::new_http("seed-peer-proxy").with_mocks(mocks);
server.start().await.map_err(|err| {
Error::Internal(format!("failed to start mock seed peer proxy: {err}"))
})?;
Ok(server)
}
async fn setup_flaky_seed_peer_proxy(
first: FlakyFirstConnection,
body: &'static str,
) -> (u16, Arc<std::sync::atomic::AtomicUsize>) {
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let connections = Arc::new(AtomicUsize::new(0));
let counter = connections.clone();
tokio::spawn(async move {
loop {
let Ok((mut stream, _)) = listener.accept().await else {
return;
};
let flaky = counter.fetch_add(1, Ordering::SeqCst) == 0;
tokio::spawn(async move {
let mut request = Vec::new();
let mut byte = [0u8; 1];
while !request.ends_with(b"\r\n\r\n") {
match stream.read(&mut byte).await {
Ok(1) => request.push(byte[0]),
_ => return,
}
}
if flaky && first == FlakyFirstConnection::Hang {
tokio::time::sleep(Duration::from_secs(60)).await;
return;
}
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = stream.write_all(head.as_bytes()).await;
let served = if flaky { &body[..body.len() / 2] } else { body };
let _ = stream.write_all(served.as_bytes()).await;
let _ = stream.shutdown().await;
});
}
});
(port, connections)
}
#[derive(Debug, Clone, Copy, PartialEq)]
enum FlakyFirstConnection {
Hang,
TruncateBody,
}
fn flaky_test_cases() -> Vec<(FlakyFirstConnection, Duration, Option<Duration>)> {
let timeout = Duration::from_millis(300);
vec![
(
FlakyFirstConnection::TruncateBody,
DEFAULT_REQUEST_TIMEOUT,
None,
),
(FlakyFirstConnection::Hang, timeout, Some(timeout)),
]
}
fn assert_flaky_retry(
first: FlakyFirstConnection,
buf: &BytesMut,
connections: usize,
elapsed: Duration,
expected_wait: Option<Duration>,
) {
assert_eq!(&buf[..], b"prefix:hello dragonfly", "first: {first:?}");
assert_eq!(connections, 2, "first: {first:?}");
if let Some(expected_wait) = expected_wait {
assert!(
(expected_wait..expected_wait * 3).contains(&elapsed),
"first: {first:?}, elapsed: {elapsed:?}"
);
}
}
#[cfg(feature = "preheat")]
fn image_index_entry(digest: &str, platform: Option<(Os, Arch)>) -> ImageIndexEntry {
ImageIndexEntry {
media_type: IMAGE_MANIFEST_MEDIA_TYPE.to_string(),
digest: digest.to_string(),
size: 0,
platform: platform.map(|(os, architecture)| oci_client::manifest::Platform {
architecture,
os,
os_version: None,
os_features: None,
variant: None,
features: None,
}),
annotations: None,
artifact_type: None,
}
}
#[tokio::test]
async fn new_success() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let result = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().retry.max_retries, 1);
}
#[tokio::test]
async fn new_invalid_params() {
let test_cases = vec![
("", None, None),
("http://0.0.0.0:4000", Some(11), None),
("http://0.0.0.0:4000", None, Some(Duration::from_secs(0))),
("http://0.0.0.0:4000", None, Some(Duration::from_secs(601))),
];
for (endpoint, max_retries, health_check_interval) in test_cases {
let mut builder = Proxy::builder().scheduler_endpoint(endpoint.to_string());
if let Some(max_retries) = max_retries {
builder = builder.max_retries(max_retries);
}
if let Some(health_check_interval) = health_check_interval {
builder = builder.health_check_interval(health_check_interval);
}
let result = builder.build().await;
assert!(
matches!(result, Err(Error::InvalidArgument(_))),
"endpoint: {endpoint}, max_retries: {max_retries:?}, health_check_interval: {health_check_interval:?}"
);
}
}
#[tokio::test]
async fn preheat_no_available_seed_peers() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatRequest {
url: "http://example.com/payload.txt".to_string(),
tag: Some("preheat".to_string()),
application: Some("dfctl".to_string()),
replicas: 1,
..Default::default()
};
let result = proxy.preheat(&request).await;
assert!(
matches!(result, Err(Error::Internal(message)) if message.contains("failed to select seed peers"))
);
}
#[tokio::test]
async fn preheat_succeeds_with_seed_peer() {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/dfdaemon.v2.DfdaemonUpload/DownloadTask");
then.pb_stream(vec![
DownloadTaskResponse {
host_id: "seed-peer-1".to_string(),
task_id: "task-1".to_string(),
peer_id: "peer-1".to_string(),
..Default::default()
},
DownloadTaskResponse {
host_id: "seed-peer-1".to_string(),
task_id: "task-1".to_string(),
peer_id: "peer-1".to_string(),
..Default::default()
},
]);
});
let mock_seed_peer = setup_mock_seed_peer(mocks).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
0,
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatRequest {
url: "http://example.com/payload.txt".to_string(),
tag: Some("preheat".to_string()),
application: Some("dfctl".to_string()),
replicas: 1,
..Default::default()
};
let result = proxy.preheat(&request).await;
assert!(result.is_ok(), "preheat should succeed: {:?}", result.err());
}
#[tokio::test]
async fn preheat_insufficient_seed_peers() {
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
0,
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatRequest {
url: "http://example.com/payload.txt".to_string(),
..Default::default()
};
let result = proxy.preheat(&request).await;
assert!(
matches!(result, Err(Error::Internal(message)) if message.contains("insufficient seed peers"))
);
}
#[tokio::test]
async fn get_streams_the_body() {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.text("hello dragonfly");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
mock_proxy.port().unwrap(),
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 1,
..Default::default()
};
let response = proxy.get(&request).await.unwrap();
assert!(response.success);
assert_eq!(response.status_code, Some(reqwest::StatusCode::OK));
let mut body = response.body.unwrap();
let mut content = Vec::new();
while let Some(chunk) = body.try_next().await.unwrap() {
content.extend_from_slice(&chunk);
}
assert_eq!(content, b"hello dragonfly");
}
#[tokio::test]
async fn get_into_fills_the_buffer() {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.text("hello dragonfly");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
mock_proxy.port().unwrap(),
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 1,
..Default::default()
};
let mut buf = BytesMut::new();
let response = proxy.get_into(&request, &mut buf).await.unwrap();
assert!(response.success);
assert_eq!(response.status_code, Some(reqwest::StatusCode::OK));
assert!(response.body.is_none());
assert_eq!(&buf[..], b"hello dragonfly");
}
#[tokio::test]
async fn get_scatters_across_replicas() {
let mut bad_mocks = MockSet::new();
bad_mocks.mock(|when, then| {
when.get().path("/file.txt");
then.status(reqwest::StatusCode::INTERNAL_SERVER_ERROR)
.text("boom");
});
let bad_proxy = setup_mock_seed_peer_proxy(bad_mocks).await.unwrap();
let mut good_mocks = MockSet::new();
good_mocks.mock(|when, then| {
when.get().path("/file.txt");
then.text("ok");
});
let good_proxy = setup_mock_seed_peer_proxy(good_mocks).await.unwrap();
let seed_peer_1 = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let seed_peer_2 = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![
create_seed_peer_host(
"seed-peer-1",
seed_peer_1.port().unwrap(),
bad_proxy.port().unwrap(),
),
create_seed_peer_host(
"seed-peer-2",
seed_peer_2.port().unwrap(),
good_proxy.port().unwrap(),
),
])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.backoff(
ExponentialBuilder::new()
.with_min_delay(Duration::from_millis(1))
.with_max_delay(Duration::from_millis(2))
.with_jitter(),
)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let mut buf = BytesMut::new();
let response = proxy.get_into(&request, &mut buf).await.unwrap();
assert!(response.success);
assert_eq!(&buf[..], b"ok");
assert_eq!(
good_proxy
.mocks()
.iter()
.map(|mock| mock.match_count())
.sum::<usize>(),
1
);
assert!(
bad_proxy
.mocks()
.iter()
.map(|mock| mock.match_count())
.sum::<usize>()
<= 1
);
}
#[tokio::test]
async fn get_retries_only_transient_answers() {
let test_cases = vec![
(reqwest::StatusCode::SERVICE_UNAVAILABLE, "backend", 2, 3),
(reqwest::StatusCode::SERVICE_UNAVAILABLE, "proxy", 2, 3),
(reqwest::StatusCode::INTERNAL_SERVER_ERROR, "backend", 1, 2),
(reqwest::StatusCode::INTERNAL_SERVER_ERROR, "proxy", 1, 2),
(reqwest::StatusCode::REQUEST_TIMEOUT, "backend", 1, 2),
(reqwest::StatusCode::REQUEST_TIMEOUT, "proxy", 1, 2),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "backend", 3, 4),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "proxy", 3, 4),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "proxy", 0, 1),
(reqwest::StatusCode::UNAUTHORIZED, "backend", 3, 1),
(reqwest::StatusCode::UNAUTHORIZED, "proxy", 3, 1),
(reqwest::StatusCode::FORBIDDEN, "backend", 3, 1),
(reqwest::StatusCode::FORBIDDEN, "proxy", 3, 1),
(reqwest::StatusCode::NOT_FOUND, "backend", 3, 1),
(reqwest::StatusCode::NOT_FOUND, "proxy", 3, 1),
];
for (status, error_type, max_retries, expected_calls) in test_cases {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.status(status)
.headers([("X-Dragonfly-Error-Type", error_type)])
.text("no");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
mock_proxy.port().unwrap(),
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.max_retries(max_retries)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 1,
..Default::default()
};
let err = proxy.get(&request).await.err().unwrap();
assert!(
err.to_string()
.contains(&format!("status_code: Some({status:?})")),
"status: {status}, error_type: {error_type}, error: {err}"
);
assert_eq!(
mock_proxy
.mocks()
.iter()
.map(|mock| mock.match_count())
.sum::<usize>(),
expected_calls,
"status: {status}"
);
}
}
#[tokio::test]
async fn get_into_retries_a_flaky_first_attempt() {
for (first, timeout, expected_wait) in flaky_test_cases() {
let (port, connections) = setup_flaky_seed_peer_proxy(first, "hello dragonfly").await;
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
port,
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 1,
timeout,
..Default::default()
};
let start = std::time::Instant::now();
let mut buf = BytesMut::from(&b"prefix:"[..]);
let response = proxy.get_into(&request, &mut buf).await.unwrap();
assert!(response.success, "first: {first:?}");
assert_flaky_retry(
first,
&buf,
connections.load(std::sync::atomic::Ordering::SeqCst),
start.elapsed(),
expected_wait,
);
}
}
#[tokio::test]
async fn get_into_with_endpoints_retries_a_flaky_first_attempt() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
for (first, timeout, expected_wait) in flaky_test_cases() {
let (port, connections) = setup_flaky_seed_peer_proxy(first, "hello dragonfly").await;
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec![format!("http://127.0.0.1:{port}")])
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
timeout,
..Default::default()
};
let start = std::time::Instant::now();
let mut buf = BytesMut::from(&b"prefix:"[..]);
let response = proxy.get_into(&request, &mut buf).await.unwrap();
assert!(response.success, "first: {first:?}");
assert_flaky_retry(
first,
&buf,
connections.load(std::sync::atomic::Ordering::SeqCst),
start.elapsed(),
expected_wait,
);
}
}
#[tokio::test]
async fn get_with_endpoints_streams_a_truncated_body_without_retrying() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let (port, connections) =
setup_flaky_seed_peer_proxy(FlakyFirstConnection::TruncateBody, "hello dragonfly")
.await;
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec![format!("http://127.0.0.1:{port}")])
.max_retries(3)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let response = proxy.get(&request).await.unwrap();
assert!(response.success);
let mut body = response.body.unwrap();
let mut content = Vec::new();
let err = loop {
match body.try_next().await {
Ok(Some(chunk)) => content.extend_from_slice(&chunk),
Ok(None) => panic!("truncated body must not end cleanly"),
Err(err) => break err,
}
};
assert!(matches!(err, Error::Internal(_)), "unexpected: {err:?}");
assert!(content.len() < b"hello dragonfly".len());
assert_eq!(connections.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test]
async fn get_with_endpoints_retries_only_transient_answers() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let test_cases = vec![
(reqwest::StatusCode::SERVICE_UNAVAILABLE, "backend", 2, 3),
(reqwest::StatusCode::SERVICE_UNAVAILABLE, "proxy", 2, 3),
(reqwest::StatusCode::REQUEST_TIMEOUT, "proxy", 1, 2),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "backend", 3, 4),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "proxy", 3, 4),
(reqwest::StatusCode::TOO_MANY_REQUESTS, "proxy", 0, 1),
(reqwest::StatusCode::UNAUTHORIZED, "proxy", 3, 1),
(reqwest::StatusCode::FORBIDDEN, "proxy", 3, 1),
(reqwest::StatusCode::NOT_FOUND, "backend", 3, 1),
];
for (status, error_type, max_retries, expected_calls) in test_cases {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.status(status)
.headers([("X-Dragonfly-Error-Type", error_type)])
.text("no");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec![format!(
"http://127.0.0.1:{}",
mock_proxy.port().unwrap()
)])
.max_retries(max_retries)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let err = proxy.get(&request).await.err().unwrap();
assert!(
err.to_string()
.contains(&format!("status_code: Some({status:?})")),
"status: {status}, error_type: {error_type}, error: {err}"
);
assert_eq!(
mock_proxy
.mocks()
.iter()
.map(|mock| mock.match_count())
.sum::<usize>(),
expected_calls,
"status: {status}, error_type: {error_type}"
);
}
}
#[tokio::test]
async fn get_error_type_backend() {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.status(reqwest::StatusCode::INTERNAL_SERVER_ERROR)
.headers([("X-Dragonfly-Error-Type", "backend")])
.text("boom");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let mock_seed_peer = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
mock_proxy.port().unwrap(),
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 1,
..Default::default()
};
let result = proxy.get(&request).await;
assert!(
matches!(result, Err(Error::BackendError(err)) if err.message.as_deref() == Some("boom"))
);
}
#[tokio::test]
async fn new_with_endpoints_invalid_params() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let test_cases = vec![
(vec![], 1, true),
(vec!["http://127.0.0.1:4001".to_string()], 11, true),
(vec!["://".to_string()], 1, false),
];
for (endpoints, max_retries, expected_invalid_argument) in test_cases {
let result = ProxyWithEndpoints::builder()
.endpoints(endpoints.clone())
.max_retries(max_retries)
.build()
.await;
if expected_invalid_argument {
assert!(
matches!(result, Err(Error::InvalidArgument(_))),
"endpoints: {endpoints:?}, max_retries: {max_retries}"
);
} else {
assert!(
matches!(result, Err(Error::Internal(_))),
"endpoints: {endpoints:?}, max_retries: {max_retries}"
);
}
}
}
#[tokio::test]
async fn get_into_with_endpoints() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.text("hello dragonfly");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let endpoints = vec![
"http://127.0.0.1:1".to_string(),
format!("http://127.0.0.1:{}", mock_proxy.port().unwrap()),
];
let proxy = ProxyWithEndpoints::builder()
.endpoints(endpoints)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let mut buf = BytesMut::new();
let response = proxy.get_into(&request, &mut buf).await.unwrap();
assert!(response.success);
assert_eq!(response.status_code, Some(reqwest::StatusCode::OK));
assert_eq!(&buf[..], b"hello dragonfly");
}
#[tokio::test]
async fn new_with_endpoints() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let endpoints = vec![
"http://127.0.0.1:4001".to_string(),
"http://127.0.0.1:4001".to_string(),
"http://127.0.0.1:4002".to_string(),
];
let proxy = ProxyWithEndpoints::builder()
.endpoints(endpoints.clone())
.build()
.await
.unwrap();
assert_eq!(proxy.retry.max_retries, 1);
assert_eq!(proxy.endpoints, endpoints);
assert_eq!(proxy.clients.len(), 2);
assert!(proxy.clients.contains_key("http://127.0.0.1:4001"));
assert!(proxy.clients.contains_key("http://127.0.0.1:4002"));
}
#[tokio::test]
async fn get_with_endpoints() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.text("hello dragonfly");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let endpoints = vec![
"http://127.0.0.1:1".to_string(),
format!("http://127.0.0.1:{}", mock_proxy.port().unwrap()),
];
let proxy = ProxyWithEndpoints::builder()
.endpoints(endpoints)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let response = proxy.get(&request).await.unwrap();
assert!(response.success);
assert_eq!(response.status_code, Some(reqwest::StatusCode::OK));
let mut body = response.body.unwrap();
let mut content = Vec::new();
while let Some(chunk) = body.try_next().await.unwrap() {
content.extend_from_slice(&chunk);
}
assert_eq!(content, b"hello dragonfly");
}
#[tokio::test]
async fn get_with_endpoints_all_endpoints_down() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec!["http://127.0.0.1:1".to_string()])
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let result = proxy.get(&request).await;
assert!(matches!(result, Err(Error::Internal(_))));
}
#[tokio::test]
async fn get_with_endpoints_invalid_replicas() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec!["http://127.0.0.1:4001".to_string()])
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 0,
..Default::default()
};
let result = proxy.get(&request).await;
assert!(matches!(result, Err(Error::InvalidArgument(_))));
}
#[tokio::test]
async fn get_with_endpoints_error_type_backend() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.get().path("/file.txt");
then.status(reqwest::StatusCode::INTERNAL_SERVER_ERROR)
.headers([("X-Dragonfly-Error-Type", "backend")])
.text("boom");
});
let mock_proxy = setup_mock_seed_peer_proxy(mocks).await.unwrap();
let proxy = ProxyWithEndpoints::builder()
.endpoints(vec![format!(
"http://127.0.0.1:{}",
mock_proxy.port().unwrap()
)])
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
..Default::default()
};
let result = proxy.get(&request).await;
assert!(
matches!(result, Err(Error::BackendError(err)) if err.message.as_deref() == Some("boom"))
);
}
#[tokio::test]
async fn get_invalid_replicas() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/file.txt".to_string(),
replicas: 0,
..Default::default()
};
let result = proxy.get(&request).await;
assert!(matches!(result, Err(Error::InvalidArgument(_))));
}
#[tokio::test]
async fn preheat_and_get_hit_same_seed_peers() {
let cases = [
("http://example.com/replicas-1.txt", 1, vec!["seed-peer-1"]),
(
"http://example.com/replicas-2.txt",
2,
vec!["seed-peer-1", "seed-peer-2"],
),
(
"http://example.com/replicas-3.txt",
3,
vec!["seed-peer-1", "seed-peer-2", "seed-peer-3"],
),
];
for (url, replicas, expected) in cases {
let path = url.strip_prefix("http://example.com").unwrap();
let mut hosts = Vec::new();
let mut servers = Vec::new();
let mut proxy_servers = Vec::new();
for name in ["seed-peer-1", "seed-peer-2", "seed-peer-3"] {
let mut mocks = MockSet::new();
if expected.contains(&name) {
mocks.mock(|when, then| {
when.path("/dfdaemon.v2.DfdaemonUpload/DownloadTask");
then.pb_stream(vec![DownloadTaskResponse {
host_id: name.to_string(),
task_id: "task-1".to_string(),
peer_id: "peer-1".to_string(),
..Default::default()
}]);
});
}
let seed_peer = setup_mock_seed_peer(mocks).await.unwrap();
let mut proxy_mocks = MockSet::new();
proxy_mocks.mock(|when, then| {
when.get().path(path);
then.status(reqwest::StatusCode::INTERNAL_SERVER_ERROR)
.text(name);
});
let seed_peer_proxy = setup_mock_seed_peer_proxy(proxy_mocks).await.unwrap();
hosts.push(create_seed_peer_host(
name,
seed_peer.port().unwrap(),
seed_peer_proxy.port().unwrap(),
));
servers.push(seed_peer);
proxy_servers.push((name, seed_peer_proxy));
}
let mock_scheduler = setup_mock_scheduler(hosts).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.max_retries(2)
.build()
.await
.unwrap();
let preheat_request = PreheatRequest {
url: url.to_string(),
replicas,
..Default::default()
};
proxy.preheat(&preheat_request).await.unwrap();
let request = GetRequest {
url: url.to_string(),
replicas,
..Default::default()
};
let mut buf = BytesMut::new();
let result = proxy.get_into(&request, &mut buf).await;
assert!(result.is_err());
for (name, server) in proxy_servers.iter() {
let hits: usize = server.mocks().iter().map(|mock| mock.match_count()).sum();
if expected.contains(name) {
assert!(hits >= 1, "{url}: expected seed peer {name} to be hit");
} else {
assert_eq!(hits, 0, "{url}: unexpected seed peer {name} was hit");
}
}
}
}
#[tokio::test]
async fn lookup_endpoints_returns_the_selected_seed_peers() {
let mut servers = Vec::new();
let mut endpoints = std::collections::HashMap::new();
let mut hosts = Vec::new();
for name in ["seed-peer-1", "seed-peer-2", "seed-peer-3"] {
let server = setup_mock_seed_peer(MockSet::new()).await.unwrap();
let port = server.port().unwrap();
endpoints.insert(name, format!("http://127.0.0.1:{port}"));
hosts.push(create_seed_peer_host(name, port, 0));
servers.push(server);
}
let mock_scheduler = setup_mock_scheduler(hosts).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let blob_url = "http://registry.example.com/v2/library/ubuntu/blobs/sha256:b2c366cce7e68013d5441c6326d5a3e1b12aeb5ed58564d0fd3fa089bc29cb6e";
let cases = [
(
GetRequest {
url: "https://example.com/file.txt?Expires=e1&Signature=s1&foo=bar".to_string(),
piece_length: Some(4194304),
tag: Some("tag-a".to_string()),
application: Some("app-a".to_string()),
filtered_query_params: vec!["Expires".to_string(), "Signature".to_string()],
replicas: 3,
..Default::default()
},
vec!["seed-peer-3", "seed-peer-2", "seed-peer-1"],
),
(
GetRequest {
url: "https://example.com/file.txt?Expires=e2&Signature=s2&foo=bar".to_string(),
piece_length: Some(4194304),
tag: Some("tag-a".to_string()),
application: Some("app-a".to_string()),
filtered_query_params: vec!["Expires".to_string(), "Signature".to_string()],
replicas: 3,
..Default::default()
},
vec!["seed-peer-3", "seed-peer-2", "seed-peer-1"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
..Default::default()
},
vec!["seed-peer-1", "seed-peer-2"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
replicas: 1,
..Default::default()
},
vec!["seed-peer-1"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
tag: Some("tag-a".to_string()),
..Default::default()
},
vec!["seed-peer-3", "seed-peer-2"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
tag: Some("tag-b".to_string()),
..Default::default()
},
vec!["seed-peer-2", "seed-peer-1"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
application: Some("app-a".to_string()),
..Default::default()
},
vec!["seed-peer-1", "seed-peer-3"],
),
(
GetRequest {
url: "https://example.com/file.txt".to_string(),
content_for_calculating_task_id: Some("This is a test file".to_string()),
replicas: 3,
..Default::default()
},
vec!["seed-peer-2", "seed-peer-3", "seed-peer-1"],
),
(
GetRequest {
url: blob_url.to_string(),
replicas: 3,
..Default::default()
},
vec!["seed-peer-3", "seed-peer-2", "seed-peer-1"],
),
(
GetRequest {
url: blob_url.to_string(),
enable_task_id_based_blob_digest: false,
..Default::default()
},
vec!["seed-peer-3", "seed-peer-1"],
),
];
for (request, expected_names) in cases {
let expected: Vec<String> = expected_names
.iter()
.map(|name| endpoints[name].clone())
.collect();
assert_eq!(proxy.lookup_endpoints(&request).await.unwrap(), expected);
}
}
#[tokio::test]
async fn lookup_endpoints_no_available_seed_peers() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = GetRequest {
url: "http://example.com/payload.txt".to_string(),
..Default::default()
};
let result = proxy.lookup_endpoints(&request).await;
assert!(
matches!(result, Err(Error::Internal(message)) if message.contains("failed to select seed peers"))
);
}
#[tokio::test]
async fn preheat_retries_on_the_same_seed_peer() {
let mut failing_mocks = MockSet::new();
failing_mocks.mock(|when, then| {
when.path("/dfdaemon.v2.DfdaemonUpload/DownloadTask");
then.error(StatusCode::SERVICE_UNAVAILABLE, "seed peer is busy");
});
let failing_seed_peer = setup_mock_seed_peer(failing_mocks).await.unwrap();
let mut healthy_mocks = MockSet::new();
healthy_mocks.mock(|when, then| {
when.path("/dfdaemon.v2.DfdaemonUpload/DownloadTask");
then.pb_stream(vec![DownloadTaskResponse {
host_id: "seed-peer-2".to_string(),
task_id: "task-1".to_string(),
peer_id: "peer-1".to_string(),
..Default::default()
}]);
});
let healthy_seed_peer = setup_mock_seed_peer(healthy_mocks).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![
create_seed_peer_host("seed-peer-1", failing_seed_peer.port().unwrap(), 0),
create_seed_peer_host("seed-peer-2", healthy_seed_peer.port().unwrap(), 0),
])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.max_retries(2)
.build()
.await
.unwrap();
let request = PreheatRequest {
url: "http://example.com/payload.txt".to_string(),
replicas: 2,
..Default::default()
};
let result = proxy.preheat(&request).await;
assert!(
matches!(&result, Err(Error::TonicStatus(status)) if status.code() == tonic::Code::Unavailable),
"unexpected: {result:?}"
);
assert_eq!(download_task_calls(&failing_seed_peer), 3);
assert_eq!(download_task_calls(&healthy_seed_peer), 1);
}
#[tokio::test]
async fn preheat_retries_only_transient_download_failures() {
let test_cases = vec![
(
StatusCode::INTERNAL_SERVER_ERROR,
None,
2,
tonic::Code::Internal,
),
(
StatusCode::SERVICE_UNAVAILABLE,
Some(2),
3,
tonic::Code::Unavailable,
),
(StatusCode::NOT_FOUND, Some(3), 1, tonic::Code::NotFound),
(
StatusCode::FORBIDDEN,
Some(3),
1,
tonic::Code::PermissionDenied,
),
];
for (status, max_retries, expected_calls, expected_code) in test_cases {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/dfdaemon.v2.DfdaemonUpload/DownloadTask");
then.error(status.clone(), "download failed");
});
let mock_seed_peer = setup_mock_seed_peer(mocks).await.unwrap();
let mock_scheduler = setup_mock_scheduler(vec![create_seed_peer_host(
"seed-peer-1",
mock_seed_peer.port().unwrap(),
0,
)])
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let mut builder = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.backoff(
ExponentialBuilder::new()
.with_min_delay(Duration::from_millis(1))
.with_max_delay(Duration::from_millis(2)),
);
if let Some(max_retries) = max_retries {
builder = builder.max_retries(max_retries);
}
let proxy = builder.build().await.unwrap();
let request = PreheatRequest {
url: "http://example.com/payload.txt".to_string(),
replicas: 1,
..Default::default()
};
let result = proxy.preheat(&request).await;
assert!(
matches!(&result, Err(Error::TonicStatus(got)) if got.code() == expected_code),
"status: {status:?}, unexpected: {result:?}"
);
assert_eq!(
download_task_calls(&mock_seed_peer),
expected_calls,
"status: {status:?}"
);
}
}
fn download_task_calls(seed_peer: &mocktail::server::MockServer) -> usize {
seed_peer
.mocks()
.iter()
.next()
.map(|mock| mock.match_count() / 2)
.unwrap_or_default()
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn preheat_image_invalid_reference() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatImageRequest {
image: "invalid image reference!!".to_string(),
..Default::default()
};
let result = proxy.preheat_image(&request).await;
assert!(
matches!(result, Err(Error::InvalidArgument(message)) if message.contains("invalid image reference"))
);
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn preheat_image_invalid_platform() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatImageRequest {
image: "docker.io/library/nginx:latest".to_string(),
platform: Some("linux-amd64".to_string()),
..Default::default()
};
let result = proxy.preheat_image(&request).await;
assert!(
matches!(result, Err(Error::InvalidArgument(message)) if message.contains("invalid platform format"))
);
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn preheat_image_unreachable_registry() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = PreheatImageRequest {
image: "127.0.0.1:1/library/nginx:latest".to_string(),
..Default::default()
};
let result = proxy.preheat_image(&request).await;
assert!(
matches!(result, Err(Error::Internal(message)) if message.contains("failed to pull image manifest"))
);
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_queries_seed_peers() {
use dragonfly_api::scheduler::v2::{
Image as ApiImage, Layer as ApiLayer, PeerImage as ApiPeerImage,
StatImageResponse as ApiStatImageResponse,
};
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/StatImage")
.pb(SchedulerStatImageRequest {
url: "https://example.com/v2/foo/bar/manifests/1.0".to_string(),
piece_length: Some(4194304),
tag: Some("stat".to_string()),
application: Some("dfctl".to_string()),
filtered_query_params: default_proxy_rule_filtered_query_params(),
username: Some("user".to_string()),
password: Some("pass".to_string()),
platform: Some("linux/amd64".to_string()),
timeout: Some(
prost_wkt_types::Duration::try_from(DEFAULT_REQUEST_TIMEOUT).unwrap(),
),
scope: STAT_IMAGE_SCOPE_ALL_SEED_PEERS.to_string(),
enable_task_id_based_blob_digest: true,
..Default::default()
});
then.pb(ApiStatImageResponse {
image: Some(ApiImage {
layers: vec![
ApiLayer {
url: "https://example.com/v2/foo/bar/blobs/sha256:b5f4dfca35398b36f61baa60e2bf2c242401c9d7db3de9168dcf780a2feedd2d".to_string(),
..Default::default()
},
ApiLayer {
url: "https://example.com/v2/foo/bar/blobs/sha256:150b7321c0794448817b19fab51e415ff406ac8663c4f53d64c3590454dee201".to_string(),
..Default::default()
},
],
}),
peers: vec![ApiPeerImage {
ip: "127.0.0.1".to_string(),
hostname: "seed-peer-1".to_string(),
cached_layers: vec![
ApiLayer {
url: "https://example.com/v2/foo/bar/blobs/sha256:b5f4dfca35398b36f61baa60e2bf2c242401c9d7db3de9168dcf780a2feedd2d".to_string(),
is_finished: Some(true),
},
ApiLayer {
url: "https://example.com/v2/foo/bar/blobs/sha256:150b7321c0794448817b19fab51e415ff406ac8663c4f53d64c3590454dee201".to_string(),
..Default::default()
},
],
}],
});
});
let mock_scheduler = setup_mock_scheduler_with_mocks(vec![], mocks)
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "example.com/foo/bar:1.0".to_string(),
username: Some("user".to_string()),
password: Some("pass".to_string()),
platform: Some("linux/amd64".to_string()),
piece_length: Some(4194304),
tag: Some("stat".to_string()),
application: Some("dfctl".to_string()),
..Default::default()
};
let response = proxy.stat_image(&request).await.unwrap();
assert_eq!(response.layers.len(), 2);
assert_eq!(response.peers.len(), 1);
assert_eq!(response.peers[0].ip, "127.0.0.1");
assert_eq!(response.peers[0].hostname, "seed-peer-1");
assert_eq!(response.peers[0].cached_layers.len(), 2);
assert!(response.peers[0].cached_layers[0].is_finished);
assert!(!response.peers[0].cached_layers[1].is_finished);
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_omits_empty_optional_fields() {
use dragonfly_api::scheduler::v2::StatImageResponse as ApiStatImageResponse;
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/StatImage")
.pb(SchedulerStatImageRequest {
url: "https://example.com/v2/foo/bar/manifests/1.0".to_string(),
filtered_query_params: default_proxy_rule_filtered_query_params(),
timeout: Some(
prost_wkt_types::Duration::try_from(DEFAULT_REQUEST_TIMEOUT).unwrap(),
),
scope: STAT_IMAGE_SCOPE_ALL_SEED_PEERS.to_string(),
enable_task_id_based_blob_digest: true,
..Default::default()
});
then.pb(ApiStatImageResponse::default());
});
let mock_scheduler = setup_mock_scheduler_with_mocks(vec![], mocks)
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "example.com/foo/bar:1.0".to_string(),
..Default::default()
};
let response = proxy.stat_image(&request).await.unwrap();
assert!(response.layers.is_empty());
assert!(response.peers.is_empty());
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_with_digest_reference() {
use dragonfly_api::scheduler::v2::StatImageResponse as ApiStatImageResponse;
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/StatImage")
.pb(SchedulerStatImageRequest {
url: "https://example.com/v2/foo/bar/manifests/sha256:b5f4dfca35398b36f61baa60e2bf2c242401c9d7db3de9168dcf780a2feedd2d".to_string(),
filtered_query_params: default_proxy_rule_filtered_query_params(),
timeout: Some(
prost_wkt_types::Duration::try_from(DEFAULT_REQUEST_TIMEOUT).unwrap(),
),
scope: STAT_IMAGE_SCOPE_ALL_SEED_PEERS.to_string(),
enable_task_id_based_blob_digest: true,
..Default::default()
});
then.pb(ApiStatImageResponse::default());
});
let mock_scheduler = setup_mock_scheduler_with_mocks(vec![], mocks)
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "example.com/foo/bar@sha256:b5f4dfca35398b36f61baa60e2bf2c242401c9d7db3de9168dcf780a2feedd2d"
.to_string(),
..Default::default()
};
let response = proxy.stat_image(&request).await.unwrap();
assert!(response.layers.is_empty());
assert!(response.peers.is_empty());
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_normalizes_docker_hub_reference() {
use dragonfly_api::scheduler::v2::StatImageResponse as ApiStatImageResponse;
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/StatImage")
.pb(SchedulerStatImageRequest {
url: "https://registry-1.docker.io/v2/library/nginx/manifests/latest"
.to_string(),
filtered_query_params: default_proxy_rule_filtered_query_params(),
timeout: Some(
prost_wkt_types::Duration::try_from(DEFAULT_REQUEST_TIMEOUT).unwrap(),
),
scope: STAT_IMAGE_SCOPE_ALL_SEED_PEERS.to_string(),
enable_task_id_based_blob_digest: true,
..Default::default()
});
then.pb(ApiStatImageResponse::default());
});
let mock_scheduler = setup_mock_scheduler_with_mocks(vec![], mocks)
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "nginx".to_string(),
..Default::default()
};
let response = proxy.stat_image(&request).await.unwrap();
assert!(response.layers.is_empty());
assert!(response.peers.is_empty());
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_invalid_reference() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "INVALID IMAGE".to_string(),
..Default::default()
};
let result = proxy.stat_image(&request).await;
assert!(matches!(result, Err(Error::InvalidArgument(_))));
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_scheduler_error() {
let mock_server = setup_mock_scheduler(vec![]).await.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_server.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "example.com/foo/bar:1.0".to_string(),
timeout: Duration::from_secs(5),
..Default::default()
};
let result = proxy.stat_image(&request).await;
assert!(
matches!(result, Err(Error::Internal(message)) if message.contains("failed to stat image"))
);
}
#[cfg(feature = "preheat")]
#[tokio::test]
async fn stat_image_scheduler_invalid_argument() {
let mut mocks = MockSet::new();
mocks.mock(|when, then| {
when.path("/scheduler.v2.Scheduler/StatImage");
then.unprocessable_content();
});
let mock_scheduler = setup_mock_scheduler_with_mocks(vec![], mocks)
.await
.unwrap();
let scheduler_endpoint = format!("http://0.0.0.0:{}", mock_scheduler.port().unwrap());
let proxy = Proxy::builder()
.scheduler_endpoint(scheduler_endpoint)
.build()
.await
.unwrap();
let request = StatImageRequest {
image: "example.com/foo/bar:1.0".to_string(),
..Default::default()
};
let result = proxy.stat_image(&request).await;
assert!(
matches!(result, Err(Error::InvalidArgument(message)) if message.contains("failed to stat image"))
);
}
#[cfg(feature = "preheat")]
#[test]
fn build_blob_url_uses_https_by_default() {
let url = Proxy::build_blob_url("registry.example.com", "library/nginx", "sha256:abcdef");
assert_eq!(
url,
"https://registry.example.com/v2/library/nginx/blobs/sha256:abcdef"
);
}
#[cfg(feature = "preheat")]
#[test]
fn build_manifest_url_uses_https_by_default() {
let url = Proxy::build_manifest_url("registry.example.com", "library/nginx", "latest");
assert_eq!(
url,
"https://registry.example.com/v2/library/nginx/manifests/latest"
);
}
#[cfg(feature = "preheat")]
#[test]
fn resolve_registry_maps_docker_hub() {
let test_cases = vec![
("nginx", "registry-1.docker.io"),
("example.com/foo/bar:1.0", "example.com"),
];
for (image, expected) in test_cases {
let reference: Reference = image.parse().unwrap();
assert_eq!(Proxy::resolve_registry(&reference), expected);
}
}
#[cfg(feature = "preheat")]
#[test]
fn platform_resolver_matches_manifests() {
let manifests = vec![
image_index_entry("sha256:amd64", Some((Os::Linux, Arch::Amd64))),
image_index_entry("sha256:arm64", Some((Os::Linux, Arch::ARM64))),
image_index_entry("sha256:no-platform", None),
];
let test_cases = vec![
("linux/amd64", Ok(Some("sha256:amd64"))),
("linux/arm64", Ok(Some("sha256:arm64"))),
("windows/amd64", Ok(None)),
("linux/riscv64", Ok(None)),
("linux-amd64", Err("invalid platform format")),
("", Err("invalid platform format")),
];
for (platform, expected) in test_cases {
match expected {
Ok(expected_digest) => {
let resolver = Proxy::platform_resolver(platform).unwrap();
assert_eq!(
resolver(&manifests),
expected_digest.map(|digest| digest.to_string()),
"platform: {platform}"
);
}
Err(expected_message) => {
assert!(
matches!(Proxy::platform_resolver(platform), Err(Error::InvalidArgument(message)) if message.contains(expected_message)),
"platform: {platform}"
);
}
}
}
}
#[tokio::test]
async fn client_pool_get_or_create() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let pool = PoolBuilder::new(HTTPClientFactory {})
.capacity(10)
.idle_timeout(Duration::from_secs(600))
.build();
assert_eq!(pool.size().await, 0);
let addr = "http://proxy1.com".to_string();
let _ = pool.entry(&addr, &addr).await.unwrap();
assert_eq!(pool.size().await, 1);
let _ = pool.entry(&addr, &addr).await.unwrap();
assert_eq!(pool.size().await, 1);
let addr = "http://proxy2.com".to_string();
let _ = pool.entry(&addr, &addr).await.unwrap();
assert_eq!(pool.size().await, 2);
}
#[tokio::test]
async fn client_pool_cleanup() {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let pool = PoolBuilder::new(HTTPClientFactory {})
.capacity(10)
.idle_timeout(Duration::from_millis(10))
.build();
let addr = "http://proxy1.com".to_string();
let _ = pool.entry(&addr, &addr).await.unwrap();
assert_eq!(pool.size().await, 1);
tokio::time::sleep(Duration::from_millis(50)).await;
let addr = "http://proxy2.com".to_string();
let _ = pool.entry(&addr, &addr).await.unwrap();
assert_eq!(pool.size().await, 1);
}
}