use crate::config::{FetchPolicy, NpmConfig};
use crate::{Error, NetworkMode, Packument};
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use std::sync::Mutex;
#[derive(Debug, Serialize, Deserialize)]
struct CachedPackument {
etag: Option<String>,
last_modified: Option<String>,
fetched_at: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
max_age_secs: Option<u64>,
packument: Packument,
}
#[derive(Debug, Serialize, Deserialize)]
struct CachedFullPackument {
etag: Option<String>,
last_modified: Option<String>,
fetched_at: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
max_age_secs: Option<u64>,
packument: serde_json::Value,
}
fn cached_is_fresh(fetched_at: u64, max_age_secs: Option<u64>) -> bool {
let age = now_secs().saturating_sub(fetched_at);
let budget = max_age_secs.unwrap_or(PACKUMENT_TTL_SECS);
age < budget
}
const PACKUMENT_TTL_SECS: u64 = 1800;
fn is_retriable_status(status: reqwest::StatusCode) -> bool {
status.is_server_error() || status == reqwest::StatusCode::TOO_MANY_REQUESTS
}
const PACKUMENT_ACCEPT: &str =
"application/vnd.npm.install-v1+json; q=1.0, application/json; q=0.8, */*";
const PACKUMENT_FULL_ACCEPT: &str = "application/json; q=1.0, */*";
const AUDIT_BODY_CAP: u64 = 256 << 20;
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
fn extract_cache_headers(resp: &reqwest::Response) -> (Option<String>, Option<String>) {
let headers = resp.headers();
let grab = |name: reqwest::header::HeaderName| {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
};
(
grab(reqwest::header::ETAG),
grab(reqwest::header::LAST_MODIFIED),
)
}
fn parse_cache_control_max_age(resp: &reqwest::Response) -> Option<u64> {
let raw = resp
.headers()
.get(reqwest::header::CACHE_CONTROL)
.and_then(|v| v.to_str().ok())?;
let mut max_age = None;
let mut s_maxage = None;
let mut force_revalidate = false;
for directive in raw.split(',').map(str::trim) {
let directive_lc = directive.to_ascii_lowercase();
match directive_lc.as_str() {
"no-store" | "no-cache" | "private" => force_revalidate = true,
_ => {}
}
if let Some(val) = directive_lc.strip_prefix("s-maxage=") {
s_maxage = val.parse::<u64>().ok();
} else if let Some(val) = directive_lc.strip_prefix("max-age=") {
max_age = val.parse::<u64>().ok();
}
}
if force_revalidate {
return Some(0);
}
s_maxage.or(max_age)
}
pub struct RegistryClient {
http: reqwest::Client,
http_by_uri: BTreeMap<String, reqwest::Client>,
token_helper_cache: Mutex<BTreeMap<String, Option<String>>>,
config: NpmConfig,
network_mode: NetworkMode,
fetch_policy: FetchPolicy,
}
impl RegistryClient {
pub fn new(registry_url: &str) -> Self {
let mut config = NpmConfig {
registry: crate::config::normalize_registry_url_pub(registry_url),
..Default::default()
};
config.apply_proxy_env();
Self::from_config(config)
}
pub fn from_config(config: NpmConfig) -> Self {
Self::from_config_with_policy(config, FetchPolicy::default())
}
pub fn from_config_with_policy(config: NpmConfig, fetch_policy: FetchPolicy) -> Self {
let http = build_http_client(&config, None, &fetch_policy);
let mut http_by_uri = BTreeMap::new();
for (uri, registry) in &config.auth_by_uri {
if registry.tls.ca.is_empty()
&& registry.tls.cafile.is_none()
&& registry.tls.cert.is_none()
&& registry.tls.key.is_none()
{
continue;
}
http_by_uri.insert(
uri.clone(),
build_http_client(&config, Some(registry), &fetch_policy),
);
}
Self {
http,
http_by_uri,
token_helper_cache: Mutex::new(BTreeMap::new()),
config,
network_mode: NetworkMode::Online,
fetch_policy,
}
}
pub fn with_network_mode(mut self, mode: NetworkMode) -> Self {
self.network_mode = mode;
self
}
pub fn network_mode(&self) -> NetworkMode {
self.network_mode
}
fn registry_url_for(&self, name: &str) -> &str {
self.config.registry_for(name)
}
fn packument_url(&self, name: &str) -> (String, &str) {
let registry_url = self.registry_url_for(name);
let url = format!(
"{}/{}",
registry_url.trim_end_matches('/'),
encoded_name(name),
);
(url, registry_url)
}
fn authed_get(&self, url: &str, registry_url: &str) -> reqwest::RequestBuilder {
self.authed_request(reqwest::Method::GET, url, registry_url)
}
pub fn authed_request(
&self,
method: reqwest::Method,
url: &str,
registry_url: &str,
) -> reqwest::RequestBuilder {
self.authed(
self.http_for(registry_url).request(method, url),
registry_url,
)
}
pub fn has_resolved_auth_for(&self, registry_url: &str) -> bool {
self.registry_auth_token_for(registry_url).is_some()
|| self.config.basic_auth_for(registry_url).is_some()
|| self.config.global_auth_token.is_some()
}
fn authed(&self, req: reqwest::RequestBuilder, registry_url: &str) -> reqwest::RequestBuilder {
if let Some(token) = self.registry_auth_token_for(registry_url) {
req.bearer_auth(token)
} else if let Some(auth) = self.config.basic_auth_for(registry_url) {
req.header("Authorization", format!("Basic {auth}"))
} else if let Some(token) = self.config.global_auth_token.as_ref()
&& same_host(&self.config.registry, registry_url)
{
req.bearer_auth(token)
} else {
req
}
}
fn registry_auth_token_for(&self, registry_url: &str) -> Option<String> {
if let Some(auth) = self.config.registry_config_for(registry_url) {
if let Some(token) = auth.auth_token.as_ref() {
return Some(token.to_string());
}
if let Some(helper) = auth.token_helper.as_deref() {
return self.cached_token_helper_result(helper);
}
}
None
}
fn cached_token_helper_result(&self, helper: &str) -> Option<String> {
{
let cache = self.token_helper_cache.lock().ok()?;
if let Some(token) = cache.get(helper) {
return token.clone();
}
}
let token = crate::config::run_token_helper(helper);
if let Ok(mut cache) = self.token_helper_cache.lock() {
cache.insert(helper.to_string(), token.clone());
}
token
}
fn http_for(&self, registry_url: &str) -> &reqwest::Client {
let uri_key = crate::config::registry_uri_key_pub(registry_url);
crate::config::lookup_by_uri_prefix(&self.http_by_uri, &uri_key).unwrap_or(&self.http)
}
async fn send_with_retry_timed<F>(
&self,
build: F,
) -> Result<(reqwest::Response, std::time::Duration), reqwest::Error>
where
F: Fn() -> reqwest::RequestBuilder,
{
let started = std::time::Instant::now();
let max_attempts = self.fetch_policy.retries.saturating_add(1);
for attempt in 0..max_attempts {
let is_last = attempt + 1 >= max_attempts;
match build().send().await {
Ok(resp) => {
let status = resp.status();
if !is_retriable_status(status) || is_last {
return Ok((resp, started.elapsed()));
}
let wait = retry_after_from(&resp)
.unwrap_or_else(|| self.fetch_policy.backoff_for_attempt(attempt + 1));
drop(resp);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
status = status.as_u16(),
"retrying HTTP request after transient failure",
);
tokio::time::sleep(wait).await;
}
Err(e) => {
if is_last {
return Err(e);
}
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %e,
"retrying HTTP request after transport error",
);
tokio::time::sleep(wait).await;
}
}
}
unreachable!("retry loop exited without returning; max_attempts was {max_attempts}")
}
async fn send_metadata_with_retry<F>(
&self,
label: &str,
build: F,
) -> Result<reqwest::Response, reqwest::Error>
where
F: Fn() -> reqwest::RequestBuilder,
{
let (resp, elapsed) = self.send_with_retry_timed(build).await?;
let threshold = self.fetch_policy.warn_timeout_ms;
let elapsed_ms = elapsed.as_millis() as u64;
if threshold > 0 && elapsed_ms > threshold {
tracing::warn!(
elapsed_ms,
threshold_ms = threshold,
label,
"slow registry metadata request exceeded fetchWarnTimeoutMs",
);
}
Ok(resp)
}
fn maybe_warn_slow_metadata(&self, label: &str, started: std::time::Instant) {
let threshold = self.fetch_policy.warn_timeout_ms;
let elapsed_ms = started.elapsed().as_millis() as u64;
if threshold > 0 && elapsed_ms > threshold {
tracing::warn!(
elapsed_ms,
threshold_ms = threshold,
label,
"slow registry metadata request exceeded fetchWarnTimeoutMs",
);
}
}
async fn retry_bytes_body_read<F>(
&self,
label: &str,
cap: u64,
build: F,
) -> Result<(bytes::Bytes, std::time::Duration), Error>
where
F: Fn() -> reqwest::RequestBuilder,
{
let max_attempts = self.fetch_policy.retries.saturating_add(1);
let mut timeout_retries: u32 = 0;
for attempt in 0..max_attempts {
let is_last = attempt + 1 >= max_attempts;
match build().send().await {
Ok(resp) if is_retriable_status(resp.status()) && !is_last => {
let wait = retry_after_from(&resp)
.unwrap_or_else(|| self.fetch_policy.backoff_for_attempt(attempt + 1));
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
status = resp.status().as_u16(),
label,
"retrying HTTP request after transient failure",
);
tokio::time::sleep(wait).await;
}
Ok(resp) => {
let resp = resp.error_for_status()?;
check_body_cap(&resp, cap, label)?;
let started = std::time::Instant::now();
match read_body_capped(resp, cap, label).await {
Ok(bytes) => return Ok((bytes, started.elapsed())),
Err(err) if !is_last => {
let is_timeout = matches!(&err, Error::Http(e) if e.is_timeout());
if is_timeout && timeout_retries >= TIMEOUT_RETRY_CAP {
return Err(err);
}
if is_timeout {
timeout_retries += 1;
}
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body read error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err),
}
}
Err(err) if !is_last => {
if err.is_timeout() {
if timeout_retries >= TIMEOUT_RETRY_CAP {
return Err(Error::Http(err));
}
timeout_retries += 1;
}
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after transport error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(Error::Http(err)),
}
}
unreachable!("retry loop exited without returning; max_attempts was {max_attempts}")
}
pub async fn fetch_packument_full_cached(
&self,
name: &str,
cache_dir: &Path,
) -> Result<serde_json::Value, Error> {
let registry_url = self.config.registry_for(name).to_string();
let cache_path = packument_full_cache_path(cache_dir, name, ®istry_url)
.ok_or_else(|| Error::InvalidName(name.to_string()))?;
let cached = read_cached_full_packument(&cache_path);
let force_cache = matches!(
self.network_mode,
NetworkMode::PreferOffline | NetworkMode::Offline
);
if let Some(c) = cached.as_ref()
&& (force_cache || cached_is_fresh(c.fetched_at, c.max_age_secs))
{
return Ok(cached.unwrap().packument);
}
if self.network_mode == NetworkMode::Offline {
return Err(Error::Offline(format!("packument for {name}")));
}
let (url, registry_url) = self.packument_url(name);
let started = std::time::Instant::now();
let cached_ref = cached.as_ref();
let label = format!("packument {name}");
let max_attempts = self.fetch_policy.retries.saturating_add(1);
for attempt in 0..max_attempts {
let is_last = attempt + 1 >= max_attempts;
match {
let mut req = self
.authed_get(&url, registry_url)
.header("Accept", PACKUMENT_FULL_ACCEPT);
if let Some(c) = cached_ref {
if let Some(ref etag) = c.etag {
req = req.header("If-None-Match", etag);
}
if let Some(ref lm) = c.last_modified {
req = req.header("If-Modified-Since", lm);
}
}
req
}
.send()
.await
{
Ok(resp) if is_retriable_status(resp.status()) && !is_last => {
let wait = retry_after_from(&resp)
.unwrap_or_else(|| self.fetch_policy.backoff_for_attempt(attempt + 1));
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
status = resp.status().as_u16(),
label,
"retrying HTTP request after transient failure",
);
tokio::time::sleep(wait).await;
}
Ok(resp) => {
if resp.status() == reqwest::StatusCode::NOT_FOUND {
self.maybe_warn_slow_metadata(&label, started);
return Err(Error::NotFound(name.to_string()));
}
if resp.status() == reqwest::StatusCode::NOT_MODIFIED
&& let Some(c) = cached.as_ref()
{
let revalidated_max_age =
parse_cache_control_max_age(&resp).or(c.max_age_secs);
let to_cache = CachedFullPackument {
etag: c.etag.clone(),
last_modified: c.last_modified.clone(),
fetched_at: now_secs(),
max_age_secs: revalidated_max_age,
packument: c.packument.clone(),
};
if let Err(e) = write_cached_full_packument(&cache_path, &to_cache) {
tracing::warn!(
"failed to write packument cache {}: {e}",
cache_path.display()
);
}
self.maybe_warn_slow_metadata(&label, started);
return Ok(c.packument.clone());
}
let (etag, last_modified) = extract_cache_headers(&resp);
let max_age_secs = parse_cache_control_max_age(&resp);
let resp = resp.error_for_status()?;
check_body_cap(&resp, self.fetch_policy.packument_max_bytes, &label)?;
match parse_full_response::<serde_json::Value>(resp).await {
Ok(packument) => {
let to_cache = CachedFullPackument {
etag,
last_modified,
fetched_at: now_secs(),
max_age_secs,
packument: packument.clone(),
};
if let Err(e) = write_cached_full_packument(&cache_path, &to_cache) {
tracing::warn!(
"failed to write packument cache {}: {e}",
cache_path.display()
);
}
self.maybe_warn_slow_metadata(&label, started);
return Ok(packument);
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body decode error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err),
}
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after transport error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err.into()),
}
}
unreachable!("retry loop exited without returning; max_attempts was {max_attempts}")
}
pub async fn fetch_packument_with_time_cached(
&self,
name: &str,
cache_dir: &Path,
) -> Result<Packument, Error> {
let registry_url = self.config.registry_for(name).to_string();
let cache_path = packument_full_cache_path(cache_dir, name, ®istry_url)
.ok_or_else(|| Error::InvalidName(name.to_string()))?;
let force_cache = matches!(
self.network_mode,
NetworkMode::PreferOffline | NetworkMode::Offline
);
if let Some(packument) = read_cached_full_packument_typed(&cache_path, force_cache) {
return Ok(packument);
}
let value = self.fetch_packument_full_cached(name, cache_dir).await?;
let packument: Packument = serde_json::from_value(value)
.map_err(|e| Error::Io(std::io::Error::new(std::io::ErrorKind::InvalidData, e)))?;
Ok(packument)
}
pub async fn fetch_packument(&self, name: &str) -> Result<Packument, Error> {
if self.network_mode == NetworkMode::Offline {
return Err(Error::Offline(format!("packument for {name}")));
}
let (url, registry_url) = self.packument_url(name);
let label = format!("packument {name}");
let max_attempts = self.fetch_policy.retries.saturating_add(1);
let started = std::time::Instant::now();
for attempt in 0..max_attempts {
let is_last = attempt + 1 >= max_attempts;
match {
let req = self.authed_get(&url, registry_url);
if force_full_packument() {
req
} else {
req.header("Accept", PACKUMENT_ACCEPT)
}
}
.send()
.await
{
Ok(resp) if is_retriable_status(resp.status()) && !is_last => {
let wait = retry_after_from(&resp)
.unwrap_or_else(|| self.fetch_policy.backoff_for_attempt(attempt + 1));
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
status = resp.status().as_u16(),
label,
"retrying HTTP request after transient failure",
);
tokio::time::sleep(wait).await;
}
Ok(resp) if resp.status() == reqwest::StatusCode::NOT_FOUND => {
self.maybe_warn_slow_metadata(&label, started);
return Err(Error::NotFound(name.to_string()));
}
Ok(resp) => {
match parse_full_response::<Packument>(resp.error_for_status()?).await {
Ok(packument) => {
self.maybe_warn_slow_metadata(&label, started);
return Ok(packument);
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body decode error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err),
}
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body decode error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err.into()),
}
}
unreachable!("retry loop exited without returning; max_attempts was {max_attempts}")
}
pub async fn fetch_packument_cached(
&self,
name: &str,
cache_dir: &Path,
) -> Result<Packument, Error> {
let registry_url = self.config.registry_for(name).to_string();
let cache_path = packument_cache_path(cache_dir, name, ®istry_url)
.ok_or_else(|| Error::InvalidName(name.to_string()))?;
let cached = read_cached_packument(&cache_path);
let force_cache = matches!(
self.network_mode,
NetworkMode::PreferOffline | NetworkMode::Offline
);
if let Some(c) = cached.as_ref()
&& (force_cache || cached_is_fresh(c.fetched_at, c.max_age_secs))
{
return Ok(cached.unwrap().packument);
}
if self.network_mode == NetworkMode::Offline {
return Err(Error::Offline(format!("packument for {name}")));
}
let (url, registry_url) = self.packument_url(name);
let cached_ref = cached.as_ref();
let label = format!("packument {name}");
let max_attempts = self.fetch_policy.retries.saturating_add(1);
let started = std::time::Instant::now();
for attempt in 0..max_attempts {
let is_last = attempt + 1 >= max_attempts;
match {
let mut req = self.authed_get(&url, registry_url);
if !force_full_packument() {
req = req.header("Accept", PACKUMENT_ACCEPT);
}
if let Some(c) = cached_ref {
if let Some(ref etag) = c.etag {
req = req.header("If-None-Match", etag);
}
if let Some(ref lm) = c.last_modified {
req = req.header("If-Modified-Since", lm);
}
}
req
}
.send()
.await
{
Ok(resp) if is_retriable_status(resp.status()) && !is_last => {
let wait = retry_after_from(&resp)
.unwrap_or_else(|| self.fetch_policy.backoff_for_attempt(attempt + 1));
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
status = resp.status().as_u16(),
label,
"retrying HTTP request after transient failure",
);
tokio::time::sleep(wait).await;
}
Ok(resp) if resp.status() == reqwest::StatusCode::NOT_FOUND => {
self.maybe_warn_slow_metadata(&label, started);
return Err(Error::NotFound(name.to_string()));
}
Ok(resp)
if resp.status() == reqwest::StatusCode::NOT_MODIFIED && cached.is_some() =>
{
let c = cached.as_ref().unwrap();
let revalidated_max_age = parse_cache_control_max_age(&resp).or(c.max_age_secs);
let to_cache = CachedPackument {
etag: c.etag.clone(),
last_modified: c.last_modified.clone(),
fetched_at: now_secs(),
max_age_secs: revalidated_max_age,
packument: c.packument.clone(),
};
if let Err(e) = write_cached_packument(&cache_path, &to_cache) {
tracing::warn!(
"failed to write packument cache {}: {e}",
cache_path.display()
);
}
self.maybe_warn_slow_metadata(&label, started);
return Ok(c.packument.clone());
}
Ok(resp) => {
let (etag, last_modified) = extract_cache_headers(&resp);
let max_age_secs = parse_cache_control_max_age(&resp);
let resp = resp.error_for_status()?;
check_body_cap(&resp, self.fetch_policy.packument_max_bytes, &label)?;
match parse_full_response::<Packument>(resp).await {
Ok(packument) => {
let to_cache = CachedPackument {
etag,
last_modified,
fetched_at: now_secs(),
max_age_secs,
packument: packument.clone(),
};
if let Err(e) = write_cached_packument(&cache_path, &to_cache) {
tracing::warn!(
"failed to write packument cache {}: {e}",
cache_path.display()
);
}
self.maybe_warn_slow_metadata(&label, started);
return Ok(packument);
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body decode error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err),
}
}
Err(err) if !is_last => {
let wait = self.fetch_policy.backoff_for_attempt(attempt + 1);
tracing::warn!(
attempt = attempt + 1,
max_attempts,
backoff_ms = wait.as_millis() as u64,
error = %err,
label,
"retrying HTTP request after response body decode error",
);
tokio::time::sleep(wait).await;
}
Err(err) => return Err(err.into()),
}
}
unreachable!("retry loop exited without returning; max_attempts was {max_attempts}")
}
pub async fn fetch_advisories_bulk(
&self,
pkg_versions: &std::collections::BTreeMap<String, Vec<String>>,
) -> Result<serde_json::Value, Error> {
let registry_url = &self.config.registry;
let url = format!(
"{}/-/npm/v1/security/advisories/bulk",
registry_url.trim_end_matches('/')
);
let body = serde_json::to_vec(pkg_versions)
.map_err(|e| Error::Io(std::io::Error::new(std::io::ErrorKind::InvalidData, e)))?;
let resp = self
.authed(self.http_for(registry_url).post(&url), registry_url)
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.body(body)
.send()
.await?;
if resp.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(serde_json::Value::Object(serde_json::Map::new()));
}
let resp = resp.error_for_status()?;
check_body_cap(&resp, AUDIT_BODY_CAP, "bulk advisories")?;
let json: serde_json::Value = resp.json().await?;
Ok(json)
}
pub async fn fetch_tarball_bytes(&self, url: &str) -> Result<bytes::Bytes, Error> {
let safe_url = aube_util::url::redact_url(url);
let parsed = reqwest::Url::parse(url)
.map_err(|e| Error::Io(std::io::Error::other(format!("invalid tarball url: {e}"))))?;
match parsed.scheme() {
"https" | "http" => {}
scheme => {
return Err(Error::Io(std::io::Error::other(format!(
"tarball {safe_url}: refusing scheme {scheme:?}",
))));
}
}
if self.network_mode == NetworkMode::Offline {
return Err(Error::Offline(format!("tarball {safe_url}")));
}
let (bytes, body_elapsed) = self
.retry_bytes_body_read(url, self.fetch_policy.tarball_max_bytes, || {
self.authed_get(url, url)
.header(reqwest::header::ACCEPT_ENCODING, "identity")
})
.await?;
warn_slow_tarball(
self.fetch_policy.min_speed_kibps,
url,
bytes.len(),
body_elapsed,
);
Ok(bytes)
}
pub async fn fetch_packument_json_fresh(&self, name: &str) -> Result<serde_json::Value, Error> {
let (url, registry_url) = self.packument_url(name);
let resp = self
.send_metadata_with_retry(&format!("packument {name}"), || {
self.authed_get(&url, registry_url)
.header("Accept", PACKUMENT_FULL_ACCEPT)
})
.await?;
if resp.status() == reqwest::StatusCode::NOT_FOUND {
return Err(Error::NotFound(name.to_string()));
}
let resp = resp.error_for_status()?;
check_body_cap(&resp, self.fetch_policy.packument_max_bytes, "packument")?;
let value: serde_json::Value = resp.json().await?;
Ok(value)
}
pub async fn put_packument(
&self,
name: &str,
body: &serde_json::Value,
otp: Option<&str>,
) -> Result<serde_json::Value, Error> {
let (url, registry_url) = self.packument_url(name);
let mut req = self.authed(
self.http_for(registry_url)
.put(&url)
.header("Content-Type", "application/json")
.json(body),
registry_url,
);
if let Some(code) = otp {
req = req.header("npm-otp", code);
}
let resp = req.send().await?;
let status = resp.status();
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
return Err(Error::RegistryWrite {
status: status.as_u16(),
body,
});
}
let value: serde_json::Value = resp.json().await.unwrap_or(serde_json::Value::Null);
Ok(value)
}
pub fn invalidate_full_packument_cache(&self, name: &str, cache_dir: &Path) {
let registry_url = self.config.registry_for(name).to_string();
if let Some(path) = packument_full_cache_path(cache_dir, name, ®istry_url) {
let _ = std::fs::remove_file(&path);
}
}
pub async fn fetch_dist_tags(
&self,
name: &str,
) -> Result<std::collections::BTreeMap<String, String>, Error> {
let registry_url = self.registry_url_for(name);
let url = dist_tag_root_url(registry_url, name);
let resp = self
.send_metadata_with_retry(&format!("dist-tags {name}"), || {
self.authed_get(&url, registry_url)
})
.await?;
check_dist_tag_status(&resp, name)?;
let map: std::collections::BTreeMap<String, String> =
resp.error_for_status()?.json().await?;
Ok(map)
}
pub async fn put_dist_tag(&self, name: &str, tag: &str, version: &str) -> Result<(), Error> {
let registry_url = self.registry_url_for(name);
let url = dist_tag_url(registry_url, name, tag);
let body = serde_json::to_string(version).map_err(std::io::Error::other)?;
let req = self
.http_for(registry_url)
.put(&url)
.header("Content-Type", "application/json")
.body(body);
let resp = self.authed(req, registry_url).send().await?;
check_dist_tag_status(&resp, name)?;
resp.error_for_status()?;
Ok(())
}
pub async fn delete_dist_tag(&self, name: &str, tag: &str) -> Result<(), Error> {
let registry_url = self.registry_url_for(name);
let url = dist_tag_url(registry_url, name, tag);
let req = self.http_for(registry_url).delete(&url);
let resp = self.authed(req, registry_url).send().await?;
if resp.status() == reqwest::StatusCode::NOT_FOUND {
return Err(Error::NotFound(format!("{name}@{tag}")));
}
if matches!(
resp.status(),
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN
) {
return Err(Error::Unauthorized);
}
resp.error_for_status()?;
Ok(())
}
pub fn tarball_url(&self, name: &str, version: &str) -> String {
let registry_url = self.registry_url_for(name);
let registry = registry_url.trim_end_matches('/');
let unscoped = if let Some(rest) = name.strip_prefix('@') {
rest.split('/').nth(1).unwrap_or(rest)
} else {
name
};
format!("{registry}/{name}/-/{unscoped}-{version}.tgz")
}
}
impl Default for RegistryClient {
fn default() -> Self {
Self::new("https://registry.npmjs.org")
}
}
fn same_host(a: &str, b: &str) -> bool {
let Ok(a) = reqwest::Url::parse(a) else {
return false;
};
let Ok(b) = reqwest::Url::parse(b) else {
return false;
};
a.scheme() == b.scheme()
&& a.host_str() == b.host_str()
&& a.port_or_known_default() == b.port_or_known_default()
}
fn build_http_client(
config: &NpmConfig,
registry_config: Option<&crate::config::AuthConfig>,
fetch_policy: &FetchPolicy,
) -> reqwest::Client {
let pool_max_idle = config.max_sockets.unwrap_or(64);
let mut builder = reqwest::Client::builder()
.user_agent("aube/0.1.0")
.timeout(std::time::Duration::from_millis(fetch_policy.timeout_ms))
.pool_max_idle_per_host(pool_max_idle)
.pool_idle_timeout(std::time::Duration::from_secs(90))
.http2_keep_alive_interval(std::time::Duration::from_secs(30))
.http2_keep_alive_timeout(std::time::Duration::from_secs(20))
.http2_keep_alive_while_idle(true)
.http2_adaptive_window(true)
.http2_initial_stream_window_size(Some(16 * 1024 * 1024))
.http2_initial_connection_window_size(Some(16 * 1024 * 1024))
.http2_max_frame_size(Some(16 * 1024 * 1024 - 1))
.tcp_nodelay(true)
.tcp_keepalive(std::time::Duration::from_secs(60))
.danger_accept_invalid_certs(!config.strict_ssl)
.min_tls_version(reqwest::tls::Version::TLS_1_2)
.redirect(reqwest::redirect::Policy::custom(|attempt| {
if attempt.previous().len() >= 10 {
return attempt.error("too many redirects");
}
if let Some(prev) = attempt.previous().last()
&& prev.scheme() == "https"
&& attempt.url().scheme() != "https"
{
return attempt.stop();
}
attempt.follow()
}))
.no_proxy();
if let Some(ip) = config.local_address {
builder = builder.local_address(Some(ip));
}
let no_proxy = config
.no_proxy
.as_deref()
.and_then(reqwest::NoProxy::from_string);
if let Some(ref url) = config.https_proxy {
match reqwest::Proxy::https(url) {
Ok(mut p) => {
if let Some(ref np) = no_proxy {
p = p.no_proxy(Some(np.clone()));
}
builder = builder.proxy(p);
}
Err(e) => tracing::warn!("ignoring https-proxy {url:?}: {e}"),
}
}
if let Some(ref url) = config.http_proxy {
match reqwest::Proxy::http(url) {
Ok(mut p) => {
if let Some(ref np) = no_proxy {
p = p.no_proxy(Some(np.clone()));
}
builder = builder.proxy(p);
}
Err(e) => tracing::warn!("ignoring http-proxy {url:?}: {e}"),
}
}
if let Some(registry_config) = registry_config {
for ca in ®istry_config.tls.ca {
match reqwest::Certificate::from_pem(ca.as_bytes()) {
Ok(cert) => builder = builder.add_root_certificate(cert),
Err(e) => tracing::warn!("ignoring invalid per-registry ca: {e}"),
}
}
if let Some(cafile) = ®istry_config.tls.cafile {
match std::fs::read(cafile) {
Ok(bytes) => match reqwest::Certificate::from_pem_bundle(&bytes) {
Ok(certs) => {
for cert in certs {
builder = builder.add_root_certificate(cert);
}
}
Err(e) => tracing::warn!("ignoring invalid cafile {}: {e}", cafile.display()),
},
Err(e) => tracing::warn!("ignoring unreadable cafile {}: {e}", cafile.display()),
}
}
if let (Some(cert), Some(key)) = (®istry_config.tls.cert, ®istry_config.tls.key) {
let mut pem = Vec::with_capacity(cert.len() + key.len() + 1);
pem.extend_from_slice(cert.as_bytes());
if !cert.ends_with('\n') {
pem.push(b'\n');
}
pem.extend_from_slice(key.as_bytes());
match reqwest::Identity::from_pem(&pem) {
Ok(identity) => builder = builder.identity(identity),
Err(e) => tracing::warn!("ignoring invalid per-registry client cert/key: {e}"),
}
}
}
builder.build().expect("failed to build HTTP client")
}
fn force_full_packument() -> bool {
std::env::var("AUBE_INTERNAL_FORCE_FULL_PACKUMENT").as_deref() == Ok("1")
}
async fn read_body_capped(
mut resp: reqwest::Response,
cap: u64,
label: &str,
) -> Result<bytes::Bytes, Error> {
if cap == 0 {
return Ok(resp.bytes().await?);
}
const STREAM_INITIAL: usize = 64 * 1024;
let initial = resp
.content_length()
.map(|len| len.min(cap) as usize)
.unwrap_or(STREAM_INITIAL);
let mut buf = bytes::BytesMut::with_capacity(initial);
while let Some(chunk) = resp.chunk().await? {
if (buf.len() as u64).saturating_add(chunk.len() as u64) > cap {
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("{label}: response body exceeds cap {cap}"),
)));
}
buf.extend_from_slice(&chunk);
}
Ok(buf.freeze())
}
fn check_body_cap(resp: &reqwest::Response, cap: u64, label: &str) -> Result<(), Error> {
if cap == 0 {
return Ok(());
}
if let Some(len) = resp.content_length()
&& len > cap
{
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("{label}: response Content-Length {len} exceeds cap {cap}"),
)));
}
Ok(())
}
fn packument_cache_path(cache_dir: &Path, name: &str, registry_url: &str) -> Option<PathBuf> {
let safe_name = aube_store::validate_and_encode_name(name)?;
let origin = registry_origin_segment(registry_url);
Some(cache_dir.join(origin).join(format!("{safe_name}.json")))
}
fn registry_origin_segment(registry_url: &str) -> String {
let digest = blake3::hash(registry_url.as_bytes()).to_hex();
format!("origin-{}", &digest.as_str()[..16])
}
fn encoded_name(name: &str) -> String {
name.replace('/', "%2F")
}
fn dist_tag_root_url(registry_url: &str, name: &str) -> String {
format!(
"{}/-/package/{}/dist-tags",
registry_url.trim_end_matches('/'),
encoded_name(name),
)
}
fn dist_tag_url(registry_url: &str, name: &str, tag: &str) -> String {
format!(
"{}/-/package/{}/dist-tags/{}",
registry_url.trim_end_matches('/'),
encoded_name(name),
tag,
)
}
fn check_dist_tag_status(resp: &reqwest::Response, name: &str) -> Result<(), Error> {
match resp.status() {
reqwest::StatusCode::NOT_FOUND => Err(Error::NotFound(name.to_string())),
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN => {
Err(Error::Unauthorized)
}
_ => Ok(()),
}
}
async fn parse_full_response<T>(resp: reqwest::Response) -> Result<T, Error>
where
T: serde::de::DeserializeOwned,
{
let bytes = resp.bytes().await?;
let mut buf = bytes.to_vec();
if let Ok(v) = simd_json::serde::from_slice::<T>(&mut buf) {
return Ok(v);
}
serde_json::from_slice::<T>(&bytes)
.map_err(|e| Error::Io(std::io::Error::new(std::io::ErrorKind::InvalidData, e)))
}
fn read_cached_packument(path: &Path) -> Option<CachedPackument> {
let mut content = std::fs::read(path).ok()?;
simd_json::serde::from_slice(&mut content).ok()
}
fn write_cached_packument(path: &Path, cached: &CachedPackument) -> std::io::Result<()> {
let json = serde_json::to_vec(cached).map_err(std::io::Error::other)?;
aube_util::fs_atomic::atomic_write(path, &json)
}
fn packument_full_cache_path(cache_dir: &Path, name: &str, registry_url: &str) -> Option<PathBuf> {
let safe_name = aube_store::validate_and_encode_name(name)?;
let origin = registry_origin_segment(registry_url);
Some(cache_dir.join(origin).join(format!("{safe_name}.json")))
}
fn read_cached_full_packument(path: &Path) -> Option<CachedFullPackument> {
let mut content = std::fs::read(path).ok()?;
simd_json::serde::from_slice(&mut content).ok()
}
fn read_cached_full_packument_typed(path: &Path, force_cache: bool) -> Option<Packument> {
#[derive(Deserialize)]
struct Typed {
fetched_at: u64,
#[serde(default)]
max_age_secs: Option<u64>,
packument: Packument,
}
let mut content = std::fs::read(path).ok()?;
let typed: Typed = simd_json::serde::from_slice(&mut content).ok()?;
if !force_cache && !cached_is_fresh(typed.fetched_at, typed.max_age_secs) {
return None;
}
Some(typed.packument)
}
fn write_cached_full_packument(path: &Path, cached: &CachedFullPackument) -> std::io::Result<()> {
let json = serde_json::to_vec(cached).map_err(std::io::Error::other)?;
aube_util::fs_atomic::atomic_write(path, &json)
}
fn warn_slow_tarball(threshold_kibps: u64, url: &str, len: usize, elapsed: std::time::Duration) {
if threshold_kibps == 0 {
return;
}
if len == 0 || elapsed <= std::time::Duration::from_secs(1) {
return;
}
let elapsed_ms = elapsed.as_millis() as u64;
let kibps = ((len as u64).saturating_mul(1000)) / elapsed_ms / 1024;
if kibps < threshold_kibps {
let safe_url = aube_util::url::redact_url(url);
tracing::warn!(
kibps,
threshold_kibps,
bytes = len,
elapsed_ms,
url = %safe_url,
"slow tarball download fell below fetchMinSpeedKiBps",
);
}
}
fn retry_after_from(resp: &reqwest::Response) -> Option<std::time::Duration> {
let raw = resp
.headers()
.get(reqwest::header::RETRY_AFTER)?
.to_str()
.ok()?;
let secs: u64 = raw.trim().parse().ok()?;
Some(std::time::Duration::from_secs(
secs.min(RETRY_AFTER_CAP_SECS),
))
}
const RETRY_AFTER_CAP_SECS: u64 = 60;
const TIMEOUT_RETRY_CAP: u32 = 1;
#[cfg(test)]
mod retry_tests {
use super::*;
use crate::config::FetchPolicy;
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn client_with(server: &MockServer, policy: FetchPolicy) -> RegistryClient {
let config = NpmConfig {
registry: format!("{}/", server.uri()),
..Default::default()
};
RegistryClient::from_config_with_policy(config, policy)
}
fn make_packument_json() -> serde_json::Value {
serde_json::json!({
"name": "demo",
"versions": {},
"dist-tags": {},
})
}
#[tokio::test]
async fn retries_on_503_then_succeeds() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(503))
.up_to_n_times(2)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(200).set_body_json(make_packument_json()))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 2,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let packument = client
.fetch_packument("demo")
.await
.expect("retry recovery");
assert_eq!(packument.name, "demo");
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 3, "expected 3 attempts (2 retries)");
}
#[tokio::test]
async fn retry_exhaustion_surfaces_final_5xx() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(503))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 1,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let err = client
.fetch_packument("demo")
.await
.expect_err("exhausted retries should error");
match err {
Error::Http(inner) => assert_eq!(inner.status().map(|s| s.as_u16()), Some(503)),
other => panic!("unexpected error: {other}"),
}
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 2, "retries=1 means 2 total attempts");
}
#[tokio::test]
async fn non_retriable_4xx_does_not_retry() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/missing"))
.respond_with(ResponseTemplate::new(404))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 3,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let err = client
.fetch_packument("missing")
.await
.expect_err("404 should surface");
assert!(matches!(err, Error::NotFound(_)));
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 1, "404 must not trigger retries");
}
#[tokio::test]
async fn retry_after_header_on_429_overrides_computed_backoff() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(429).insert_header("Retry-After", "0"))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(200).set_body_json(make_packument_json()))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 2,
retry_factor: 1,
retry_min_timeout_ms: 60_000,
retry_max_timeout_ms: 60_000,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let packument = tokio::time::timeout(
std::time::Duration::from_secs(2),
client.fetch_packument("demo"),
)
.await
.expect("Retry-After should be honored, overriding the 60s default backoff")
.expect("request should succeed");
assert_eq!(packument.name, "demo");
}
#[tokio::test]
async fn retries_on_429_rate_limit() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(429))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(200).set_body_json(make_packument_json()))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 2,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let packument = client.fetch_packument("demo").await.expect("429 retry");
assert_eq!(packument.name, "demo");
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 2);
}
#[tokio::test]
async fn tarball_fetch_requests_identity_encoding() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/pkg.tgz"))
.and(header("accept-encoding", "identity"))
.respond_with(ResponseTemplate::new(200).set_body_bytes(b"tgz bytes".to_vec()))
.expect(1)
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 0,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let url = format!("{}/pkg.tgz", server.uri());
let bytes = client
.fetch_tarball_bytes(&url)
.await
.expect("tarball fetch should succeed");
assert_eq!(&bytes[..], b"tgz bytes");
}
#[tokio::test]
async fn fetch_timeout_triggers_transport_error() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(make_packument_json())
.set_delay(std::time::Duration::from_millis(500)),
)
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 50,
retries: 0,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let err = client
.fetch_packument("demo")
.await
.expect_err("timeout should surface");
match err {
Error::Http(inner) => assert!(
inner.is_timeout() || inner.is_request(),
"expected timeout-shaped reqwest error, got {inner:?}",
),
other => panic!("unexpected error: {other}"),
}
}
#[tokio::test]
async fn tarball_headers_timeout_retries_at_most_once_even_with_high_retry_budget() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/pkg.tgz"))
.respond_with(
ResponseTemplate::new(200)
.set_body_bytes(b"unused".to_vec())
.set_delay(std::time::Duration::from_millis(500)),
)
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 50,
retries: 5,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let url = format!("{}/pkg.tgz", server.uri());
let err = client
.fetch_tarball_bytes(&url)
.await
.expect_err("timeout should surface");
match err {
Error::Http(inner) => assert!(
inner.is_timeout() || inner.is_request(),
"expected timeout-shaped reqwest error, got {inner:?}",
),
other => panic!("unexpected error: {other}"),
}
let requests = server.received_requests().await.unwrap();
assert_eq!(
requests.len(),
2,
"timeouts must cap retries at 1 regardless of fetchRetries",
);
}
#[tokio::test]
async fn timeout_cap_counts_only_timeouts_not_other_retries() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/pkg.tgz"))
.respond_with(ResponseTemplate::new(503))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/pkg.tgz"))
.respond_with(
ResponseTemplate::new(200)
.set_body_bytes(b"unused".to_vec())
.set_delay(std::time::Duration::from_millis(500)),
)
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 50,
retries: 5,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let url = format!("{}/pkg.tgz", server.uri());
let _ = client
.fetch_tarball_bytes(&url)
.await
.expect_err("all attempts fail");
let requests = server.received_requests().await.unwrap();
assert_eq!(
requests.len(),
3,
"expected 1 503 + 1 initial timeout + 1 capped timeout retry; \
timeout cap must not consume non-timeout retry slots",
);
}
#[tokio::test]
async fn tarball_body_read_timeout_retries_at_most_once() {
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let count = std::sync::Arc::new(AtomicUsize::new(0));
let count_handle = count.clone();
tokio::spawn(async move {
loop {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
count_handle.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let _ = sock
.write_all(
b"HTTP/1.1 200 OK\r\n\
Content-Length: 1048576\r\n\
Content-Type: application/octet-stream\r\n\r\n",
)
.await;
let _ = sock.flush().await;
tokio::time::sleep(std::time::Duration::from_secs(2)).await;
});
}
});
let policy = FetchPolicy {
timeout_ms: 100,
retries: 5,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let config = NpmConfig {
registry: format!("http://{addr}/"),
..Default::default()
};
let client = RegistryClient::from_config_with_policy(config, policy);
let url = format!("http://{addr}/pkg.tgz");
let err = client
.fetch_tarball_bytes(&url)
.await
.expect_err("body-read timeout should surface");
assert!(
matches!(&err, Error::Http(e) if e.is_timeout() || e.is_request()),
"expected timeout-shaped error, got {err:?}",
);
assert_eq!(
count.load(Ordering::SeqCst),
2,
"body-read timeouts must cap retries at 1 regardless of fetchRetries",
);
}
#[tokio::test]
async fn warn_timeout_is_pure_observability_and_does_not_fail_request() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(make_packument_json())
.set_delay(std::time::Duration::from_millis(50)),
)
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 0,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
warn_timeout_ms: 1,
min_speed_kibps: 0,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let packument = client
.fetch_packument("demo")
.await
.expect("warn-threshold is advisory — request must still succeed");
assert_eq!(packument.name, "demo");
}
#[tokio::test]
async fn retries_on_packument_body_decode_error_then_succeeds() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(
ResponseTemplate::new(200).set_body_raw("{not valid json", "application/json"),
)
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(200).set_body_json(make_packument_json()))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 2,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let packument = client
.fetch_packument("demo")
.await
.expect("decode error should be retried");
assert_eq!(packument.name, "demo");
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 2, "expected retry after decode error");
}
#[tokio::test]
async fn full_packument_cached_retries_on_body_decode_error_then_succeeds() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(
ResponseTemplate::new(200).set_body_raw("{not valid json", "application/json"),
)
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"name": "demo",
"versions": {},
"dist-tags": {},
"time": {},
})))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 2,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let temp = tempfile::tempdir().unwrap();
let packument = client
.fetch_packument_full_cached("demo", temp.path())
.await
.expect("decode error should be retried on full packument path");
assert_eq!(packument["name"], "demo");
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 2, "expected retry after decode error");
}
#[tokio::test]
async fn body_decode_retry_does_not_multiply_total_attempt_count() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(
ResponseTemplate::new(200).set_body_raw("{not valid json", "application/json"),
)
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/demo"))
.respond_with(ResponseTemplate::new(503))
.mount(&server)
.await;
let policy = FetchPolicy {
timeout_ms: 5_000,
retries: 1,
retry_factor: 1,
retry_min_timeout_ms: 1,
retry_max_timeout_ms: 1,
..FetchPolicy::default()
};
let client = client_with(&server, policy);
let err = client
.fetch_packument("demo")
.await
.expect_err("retry budget should be exhausted after two total attempts");
match err {
Error::Http(inner) => assert_eq!(inner.status().map(|s| s.as_u16()), Some(503)),
other => panic!("unexpected error: {other}"),
}
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 2, "expected total attempts to stay capped");
}
#[tokio::test]
async fn scoped_packument_request_is_url_encoded() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"name": "@scope/pkg",
"versions": {},
"dist-tags": {},
})))
.mount(&server)
.await;
let client = client_with(&server, FetchPolicy::default());
let packument = client
.fetch_packument("@scope/pkg")
.await
.expect("scoped packument fetch must succeed");
assert_eq!(packument.name, "@scope/pkg");
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 1);
let raw = requests[0].url.as_str();
assert!(
raw.contains("/@scope%2Fpkg"),
"expected %2F-encoded scope separator, got {raw}"
);
let accept = requests[0]
.headers
.get("accept")
.and_then(|v| v.to_str().ok())
.unwrap_or_default();
assert_eq!(
accept, "application/vnd.npm.install-v1+json; q=1.0, application/json; q=0.8, */*",
"corgi Accept header must include JSON and */* fallbacks",
);
}
}
#[cfg(test)]
mod slow_tarball_tests {
use super::warn_slow_tarball;
use std::time::Duration;
#[test]
fn zero_threshold_disables_warning() {
warn_slow_tarball(
0,
"https://example.com/pkg.tgz",
1024,
Duration::from_secs(10),
);
}
#[test]
fn sub_second_transfer_skipped_to_avoid_handshake_noise() {
warn_slow_tarball(
50,
"https://example.com/quick.tgz",
2048,
Duration::from_millis(500),
);
}
#[test]
fn exactly_one_second_skipped() {
warn_slow_tarball(
50,
"https://example.com/boundary.tgz",
10_240,
Duration::from_secs(1),
);
}
#[test]
fn zero_elapsed_skipped_to_avoid_division_by_zero() {
warn_slow_tarball(50, "https://example.com/fast.tgz", 10_240, Duration::ZERO);
}
#[test]
fn fast_download_does_not_warn() {
warn_slow_tarball(
50,
"https://example.com/pkg.tgz",
10 * 1024 * 1024,
Duration::from_secs(2),
);
}
#[test]
fn slow_download_triggers_warning_path() {
warn_slow_tarball(
50,
"https://example.com/slow.tgz",
10_240,
Duration::from_secs(2),
);
}
}