pub mod artifactory;
pub mod nexus;
use super::{Result, SourceKind, SourceRegistry};
use std::time::Duration;
const MAX_ATTEMPTS: u32 = 5;
const MAX_BACKOFF: Duration = Duration::from_secs(30);
pub fn build_source(
kind: SourceKind,
url: &str,
client: reqwest::Client,
auth: Option<String>,
allow_private: bool,
) -> Result<Box<dyn SourceRegistry>> {
let base = url.trim_end_matches('/').to_string();
if base.is_empty() {
return Err("--url must not be empty".to_string());
}
let http = SourceHttp {
client,
base,
auth,
allow_private,
};
match kind {
SourceKind::Artifactory => Ok(Box::new(artifactory::Artifactory::new(http))),
SourceKind::Nexus => Ok(Box::new(nexus::Nexus::new(http))),
}
}
pub(crate) struct SourceHttp {
pub(crate) client: reqwest::Client,
pub(crate) base: String,
pub(crate) auth: Option<String>,
pub(crate) allow_private: bool,
}
impl SourceHttp {
pub(crate) fn url(&self, path: &str) -> String {
format!("{}{}", self.base, path)
}
fn authed(&self, rb: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
match &self.auth {
Some(creds) => rb.header(
reqwest::header::AUTHORIZATION,
crate::config::basic_auth_header(creds),
),
None => rb,
}
}
fn authed_same_origin(
&self,
url: &str,
rb: reqwest::RequestBuilder,
) -> reqwest::RequestBuilder {
if same_origin(&self.base, url) {
self.authed(rb)
} else {
rb
}
}
pub(crate) async fn get(&self, path: &str) -> Result<reqwest::Response> {
let full = self.url(path);
self.send_with_retry(&full, "GET", || self.authed(self.client.get(&full)))
.await
}
pub(crate) async fn get_absolute(&self, url: &str) -> Result<reqwest::Response> {
if super::http::is_blocked_url_host(url, self.allow_private) {
return Err(format!(
"SSRF guard: source-supplied URL {} resolves to a blocked (loopback/private/metadata) IP literal",
super::http::redact_url(url)
));
}
self.send_with_retry(url, "GET", || {
self.authed_same_origin(url, self.client.get(url))
})
.await
}
pub(crate) async fn post_text(&self, path: &str, body: String) -> Result<reqwest::Response> {
let full = self.url(path);
self.send_with_retry(&full, "POST", || {
self.authed(
self.client
.post(&full)
.header(reqwest::header::CONTENT_TYPE, "text/plain")
.body(body.clone()),
)
})
.await
}
async fn send_with_retry<F>(
&self,
url: &str,
method: &str,
make: F,
) -> Result<reqwest::Response>
where
F: Fn() -> reqwest::RequestBuilder,
{
let safe = super::http::redact_url(url);
let mut attempt: u32 = 0;
loop {
attempt += 1;
match make().send().await {
Ok(resp) => {
let status = resp.status();
let retryable = status == reqwest::StatusCode::TOO_MANY_REQUESTS
|| status.is_server_error();
if retryable && attempt < MAX_ATTEMPTS {
let delay = retry_after(&resp).unwrap_or_else(|| backoff(attempt));
tracing::warn!(url = %safe, status = status.as_u16(), attempt, delay_ms = delay.as_millis() as u64, "import source retry");
tokio::time::sleep(delay).await;
continue;
}
return Ok(resp);
}
Err(e) => {
let transient =
!e.is_redirect() && (e.is_timeout() || e.is_connect() || e.is_request());
if transient && attempt < MAX_ATTEMPTS {
let delay = backoff(attempt);
tracing::warn!(url = %safe, attempt, delay_ms = delay.as_millis() as u64, "import source transport retry");
tokio::time::sleep(delay).await;
continue;
}
return Err(format!(
"{method} {safe} failed (timeout={}, connect={}, redirect={})",
e.is_timeout(),
e.is_connect(),
e.is_redirect()
));
}
}
}
}
}
fn retry_after(resp: &reqwest::Response) -> Option<Duration> {
let raw = resp
.headers()
.get(reqwest::header::RETRY_AFTER)?
.to_str()
.ok()?
.trim()
.to_string();
match raw.parse::<u64>() {
Ok(secs) => Some(Duration::from_secs(secs).min(MAX_BACKOFF)),
Err(_) => Some(MAX_BACKOFF),
}
}
fn backoff(attempt: u32) -> Duration {
let ms = 200u64.saturating_mul(1u64 << attempt.min(8).saturating_sub(1));
Duration::from_millis(ms).min(MAX_BACKOFF)
}
fn same_origin(base: &str, url: &str) -> bool {
let origin = |s: &str| {
reqwest::Url::parse(s).ok().map(|u| {
(
u.scheme().to_string(),
u.host_str().map(str::to_string),
u.port_or_known_default(),
)
})
};
match (origin(base), origin(url)) {
(Some(a), Some(b)) => a == b,
_ => false,
}
}
pub(crate) fn join_path(path: &str, name: &str) -> String {
let path = path.trim_matches('/');
if path.is_empty() || path == "." {
name.to_string()
} else {
format!("{path}/{name}")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backoff_is_capped_and_monotonic() {
assert_eq!(backoff(1), Duration::from_millis(200));
assert_eq!(backoff(2), Duration::from_millis(400));
assert_eq!(backoff(3), Duration::from_millis(800));
assert!(backoff(20) <= MAX_BACKOFF);
}
#[test]
fn join_path_handles_root_and_nested() {
assert_eq!(join_path(".", "foo-1.0.jar"), "foo-1.0.jar");
assert_eq!(join_path("", "foo.txt"), "foo.txt");
assert_eq!(
join_path("com/example/foo/1.0", "foo-1.0.jar"),
"com/example/foo/1.0/foo-1.0.jar"
);
assert_eq!(join_path("/a/b/", "c.bin"), "a/b/c.bin");
}
#[test]
fn same_origin_compares_scheme_host_port() {
assert!(same_origin(
"https://art.example.com",
"https://art.example.com/repo/x.jar"
));
assert!(same_origin(
"https://art.example.com:443",
"https://art.example.com/x"
)); assert!(!same_origin(
"https://art.example.com",
"https://evil.example.com/x"
)); assert!(!same_origin(
"https://art.example.com",
"http://art.example.com/x"
)); assert!(!same_origin(
"https://art.example.com",
"https://art.example.com:8443/x"
)); assert!(!same_origin("https://art.example.com", "not a url")); }
#[tokio::test]
async fn get_absolute_rejects_source_supplied_metadata_url() {
let http = SourceHttp {
client: super::super::http::build_import_client(
&crate::config::TlsConfig::default(),
Duration::from_secs(5),
Duration::from_secs(5),
false,
)
.unwrap(),
base: "https://art.example.com".to_string(),
auth: Some("u:p".to_string()),
allow_private: false,
};
for bad in [
"http://169.254.169.254/latest/meta-data/",
"http://127.0.0.1:8081/x",
"http://[::1]/x",
"http://10.0.0.5/x",
] {
let err = http.get_absolute(bad).await.unwrap_err();
assert!(
err.contains("SSRF"),
"expected SSRF rejection for {bad}, got: {err}"
);
}
}
}