use async_trait::async_trait;
use axum::body::Bytes;
use futures::TryStreamExt;
use object_store::aws::AmazonS3Builder;
use object_store::gcp::GoogleCloudStorageBuilder;
use object_store::path::Path;
use object_store::{ObjectStore, ObjectStoreExt, PutPayload, WriteMultipart};
use std::pin::Pin;
use tokio::io::{AsyncRead, AsyncReadExt};
use super::{FileMeta, Result, StorageBackend, StorageError};
const STREAM_ATTEMPT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(3600);
fn streaming_client_options() -> object_store::ClientOptions {
object_store::ClientOptions::new().with_timeout(STREAM_ATTEMPT_TIMEOUT)
}
fn streaming_retry_config() -> object_store::RetryConfig {
object_store::RetryConfig {
retry_timeout: STREAM_ATTEMPT_TIMEOUT.saturating_mul(2),
..Default::default()
}
}
pub struct ObjectStorage {
store: Box<dyn ObjectStore>,
name: &'static str,
cached_total_size: std::sync::atomic::AtomicU64,
size_cache_initialized: std::sync::atomic::AtomicBool,
cached_reachable: std::sync::atomic::AtomicBool,
last_refresh_unix: std::sync::atomic::AtomicU64,
}
const HEALTH_MAX_STALE_SECS: i64 = 150;
impl ObjectStorage {
pub fn new(
s3_url: &str,
bucket: &str,
region: &str,
access_key: Option<&str>,
secret_key: Option<&str>,
virtual_hosted: bool,
) -> Self {
let url = s3_url.trim_end_matches('/');
let allow_http = url.starts_with("http://");
let mut builder = AmazonS3Builder::new()
.with_endpoint(url)
.with_bucket_name(bucket)
.with_region(region)
.with_client_options(streaming_client_options().with_allow_http(allow_http))
.with_virtual_hosted_style_request(virtual_hosted)
.with_retry(streaming_retry_config());
match (access_key, secret_key) {
(Some(ak), Some(sk)) => {
builder = builder.with_access_key_id(ak).with_secret_access_key(sk);
}
_ => {
builder = builder.with_skip_signature(true);
}
}
let store = builder.build().expect("Failed to build S3 client");
Self {
store: Box::new(store),
name: "s3",
cached_total_size: std::sync::atomic::AtomicU64::new(0),
size_cache_initialized: std::sync::atomic::AtomicBool::new(false),
cached_reachable: std::sync::atomic::AtomicBool::new(false),
last_refresh_unix: std::sync::atomic::AtomicU64::new(0),
}
}
pub fn new_gcs(
bucket: &str,
service_account_path: Option<&str>,
base_url: Option<&str>,
) -> Self {
let mut builder = GoogleCloudStorageBuilder::from_env()
.with_bucket_name(bucket)
.with_retry(streaming_retry_config());
if std::env::var_os("GOOGLE_TIMEOUT").is_none() {
builder = builder.with_config(
object_store::gcp::GoogleConfigKey::Client(object_store::ClientConfigKey::Timeout),
format!("{}s", STREAM_ATTEMPT_TIMEOUT.as_secs()),
);
}
if let Some(path) = service_account_path {
builder = builder.with_service_account_path(path);
}
if let Some(url) = base_url {
let allow_http = url.starts_with("http://");
builder = builder
.with_base_url(url)
.with_client_options(streaming_client_options().with_allow_http(allow_http));
if allow_http {
builder = builder.with_skip_signature(true);
}
}
let store = builder.build().expect("Failed to build GCS client");
Self {
store: Box::new(store),
name: "gcs",
cached_total_size: std::sync::atomic::AtomicU64::new(0),
size_cache_initialized: std::sync::atomic::AtomicBool::new(false),
cached_reachable: std::sync::atomic::AtomicBool::new(false),
last_refresh_unix: std::sync::atomic::AtomicU64::new(0),
}
}
}
fn encode_object_key(key: &str) -> String {
key.replace('@', "%40")
}
fn encode_object_key_legacy(key: &str) -> String {
key.replace('@', "_at_")
}
fn decode_object_key(key: &str) -> String {
key.replace("%2540", "@").replace("%40", "@")
}
fn map_err(e: object_store::Error) -> StorageError {
match e {
object_store::Error::NotFound { .. } => StorageError::NotFound,
other => StorageError::Network(other.to_string()),
}
}
#[async_trait]
impl StorageBackend for ObjectStorage {
async fn put(&self, key: &str, data: &[u8]) -> Result<()> {
let encoded = encode_object_key(key);
let path = Path::from(encoded);
let payload = PutPayload::from(data.to_vec());
self.store.put(&path, payload).await.map_err(map_err)?;
Ok(())
}
async fn get(&self, key: &str) -> Result<Bytes> {
let encoded = encode_object_key(key);
let path = Path::from(encoded);
match self.store.get(&path).await {
Ok(result) => {
let bytes = result.bytes().await.map_err(map_err)?;
Ok(bytes)
}
Err(object_store::Error::NotFound { .. }) if key.contains('@') => {
let legacy_path = Path::from(encode_object_key_legacy(key));
let result = self.store.get(&legacy_path).await.map_err(map_err)?;
let bytes = result.bytes().await.map_err(map_err)?;
Ok(bytes)
}
Err(e) => Err(map_err(e)),
}
}
async fn delete(&self, key: &str) -> Result<()> {
let encoded = encode_object_key(key);
let path = Path::from(encoded);
self.store.delete(&path).await.map_err(map_err)?;
Ok(())
}
async fn list(&self, prefix: &str) -> Result<Vec<String>> {
let encoded = encode_object_key(prefix);
let prefix_path = Path::from(encoded);
let list_prefix = if prefix.is_empty() {
None
} else {
Some(&prefix_path)
};
let objects: Vec<_> = self
.store
.list(list_prefix)
.try_collect()
.await
.map_err(|e| StorageError::Network(e.to_string()))?;
Ok(objects
.into_iter()
.map(|meta| decode_object_key(meta.location.as_ref()))
.collect())
}
async fn list_with_meta(&self, prefix: &str) -> Result<Vec<(String, FileMeta)>> {
let encoded = encode_object_key(prefix);
let prefix_path = Path::from(encoded);
let list_prefix = if prefix.is_empty() {
None
} else {
Some(&prefix_path)
};
let objects: Vec<_> = self
.store
.list(list_prefix)
.try_collect()
.await
.map_err(|e| StorageError::Network(e.to_string()))?;
Ok(objects
.into_iter()
.map(|meta| {
let modified = meta.last_modified.timestamp().try_into().unwrap_or(0u64);
(
decode_object_key(meta.location.as_ref()),
FileMeta {
size: meta.size,
modified,
},
)
})
.collect())
}
async fn stat(&self, key: &str) -> Option<FileMeta> {
let encoded = encode_object_key(key);
let path = Path::from(encoded);
let meta = match self.store.head(&path).await {
Ok(m) => m,
Err(_) if key.contains('@') => {
let legacy_path = Path::from(encode_object_key_legacy(key));
self.store.head(&legacy_path).await.ok()?
}
Err(_) => return None,
};
let modified = meta.last_modified.timestamp().try_into().unwrap_or(0u64);
Some(FileMeta {
size: meta.size,
modified,
})
}
async fn health_check(&self) -> bool {
self.cached_reachable
.load(std::sync::atomic::Ordering::Relaxed)
&& crate::cache_ttl::is_within_ttl(
self.last_refresh_unix
.load(std::sync::atomic::Ordering::Relaxed),
HEALTH_MAX_STALE_SECS,
)
}
async fn total_size(&self) -> u64 {
self.cached_total_size
.load(std::sync::atomic::Ordering::Relaxed)
}
fn backend_name(&self) -> &'static str {
self.name
}
async fn refresh_total_size(&self) {
let result: std::result::Result<Vec<_>, _> = self.store.list(None).try_collect().await;
self.cached_reachable
.store(result.is_ok(), std::sync::atomic::Ordering::Relaxed);
self.last_refresh_unix.store(
crate::cache_ttl::now_unix(),
std::sync::atomic::Ordering::Relaxed,
);
if let Ok(objects) = result {
let total: u64 = objects.iter().map(|m| m.size).sum();
self.cached_total_size
.store(total, std::sync::atomic::Ordering::Relaxed);
self.size_cache_initialized
.store(true, std::sync::atomic::Ordering::Relaxed);
}
}
async fn put_from_path(&self, key: &str, src: &std::path::Path) -> Result<()> {
let encoded = encode_object_key(key);
let s3_path = Path::from(encoded);
let mut file = tokio::fs::File::open(src).await?;
let upload = self.store.put_multipart(&s3_path).await.map_err(map_err)?;
let mut writer = WriteMultipart::new(upload);
let mut buf = vec![0u8; 8 * 1024 * 1024]; loop {
let n = file.read(&mut buf).await?;
if n == 0 {
break;
}
writer.write(&buf[..n]);
}
writer.finish().await.map_err(map_err)?;
let _ = tokio::fs::remove_file(src).await;
Ok(())
}
async fn copy(&self, src: &str, dst: &str) -> Result<()> {
let from_path = Path::from(encode_object_key(src));
let to_path = Path::from(encode_object_key(dst));
self.store.copy(&from_path, &to_path).await.map_err(map_err)
}
async fn get_reader(&self, key: &str) -> Result<(u64, Pin<Box<dyn AsyncRead + Send + Unpin>>)> {
let encoded = encode_object_key(key);
let path = Path::from(encoded);
let result = match self.store.get(&path).await {
Ok(r) => r,
Err(object_store::Error::NotFound { .. }) if key.contains('@') => {
let legacy_path = Path::from(encode_object_key_legacy(key));
self.store.get(&legacy_path).await.map_err(map_err)?
}
Err(e) => return Err(map_err(e)),
};
let size = result.meta.size;
let stream = result.into_stream().map_err(std::io::Error::other);
let reader = tokio_util::io::StreamReader::new(stream);
Ok((size as u64, Box::pin(reader)))
}
async fn get_range(
&self,
key: &str,
start: u64,
end: u64,
) -> Result<(u64, Pin<Box<dyn AsyncRead + Send + Unpin>>)> {
let make_opts = || object_store::GetOptions {
range: Some(object_store::GetRange::Bounded(start..(end + 1))),
..Default::default()
};
let path = Path::from(encode_object_key(key));
let result = match self.store.get_opts(&path, make_opts()).await {
Ok(r) => r,
Err(object_store::Error::NotFound { .. }) if key.contains('@') => {
let legacy_path = Path::from(encode_object_key_legacy(key));
self.store
.get_opts(&legacy_path, make_opts())
.await
.map_err(map_err)?
}
Err(e) => return Err(map_err(e)),
};
let size = result.meta.size;
let stream = result.into_stream().map_err(std::io::Error::other);
let reader = tokio_util::io::StreamReader::new(stream);
Ok((size as u64, Box::pin(reader)))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_backend_name() {
let storage = ObjectStorage::new(
"http://localhost:9000",
"test-bucket",
"us-east-1",
Some("access"),
Some("secret"),
false,
);
assert_eq!(storage.backend_name(), "s3");
}
#[tokio::test]
async fn test_health_check_cached_not_live() {
let storage = ObjectStorage::new("http://127.0.0.1:1", "b", "r", None, None, false);
assert!(!storage.health_check().await);
storage.refresh_total_size().await;
assert!(!storage.health_check().await);
}
#[tokio::test]
async fn test_health_check_expires_when_refresh_stalls() {
use std::sync::atomic::Ordering::Relaxed;
let storage = ObjectStorage::new("http://127.0.0.1:1", "b", "r", None, None, false);
let now = crate::cache_ttl::now_unix();
storage.cached_reachable.store(true, Relaxed);
storage
.last_refresh_unix
.store(now.saturating_sub(10_000), Relaxed);
assert!(
!storage.health_check().await,
"stale reachability must not keep readiness up"
);
storage.last_refresh_unix.store(now, Relaxed);
assert!(storage.health_check().await);
}
#[test]
fn decode_reverses_double_encoded_at() {
assert_eq!(
decode_object_key("npm/%2540scope/pkg/metadata.json"),
"npm/@scope/pkg/metadata.json"
);
assert_eq!(decode_object_key("npm/%40scope/x"), "npm/@scope/x");
assert_eq!(
decode_object_key("cargo/look_at_this/x"),
"cargo/look_at_this/x"
);
}
#[tokio::test]
async fn scoped_key_lists_and_gets_through_path_encoding() {
let storage = ObjectStorage {
store: Box::new(object_store::memory::InMemory::new()),
name: "s3",
cached_total_size: std::sync::atomic::AtomicU64::new(0),
size_cache_initialized: std::sync::atomic::AtomicBool::new(false),
cached_reachable: std::sync::atomic::AtomicBool::new(true),
last_refresh_unix: std::sync::atomic::AtomicU64::new(0),
};
let key = "npm/@scope/pkg/versions/1.0.0.json";
storage.put(key, br#"{"version":"1.0.0"}"#).await.unwrap();
let plain = "npm/plainpkg/versions/1.0.0.json";
storage.put(plain, br#"{"version":"1.0.0"}"#).await.unwrap();
let listed = storage.list("npm/@scope/pkg/versions/").await.unwrap();
assert_eq!(
listed,
vec![key.to_string()],
"list must decode %2540 back to @ so scan-regenerate can read it"
);
let got = storage
.get(&listed[0])
.await
.expect("get on the listed key must succeed");
assert_eq!(&got[..], br#"{"version":"1.0.0"}"#);
let listed_plain = storage.list("npm/plainpkg/versions/").await.unwrap();
assert_eq!(listed_plain, vec![plain.to_string()]);
assert!(storage.get(&listed_plain[0]).await.is_ok());
}
#[test]
fn test_s3_storage_creation_anonymous() {
let storage = ObjectStorage::new(
"http://localhost:9000",
"test-bucket",
"us-east-1",
None,
None,
false,
);
assert_eq!(storage.backend_name(), "s3");
}
#[test]
fn test_gcs_storage_creation_emulator() {
let storage = ObjectStorage::new_gcs("test-bucket", None, Some("http://localhost:4443"));
assert_eq!(storage.backend_name(), "gcs");
}
#[test]
fn test_gcs_storage_creation_default_endpoint() {
let storage = ObjectStorage::new_gcs("test-bucket", None, None);
assert_eq!(storage.backend_name(), "gcs");
}
const EMPTY_LIST_XML: &str = r#"<?xml version="1.0" encoding="UTF-8"?><ListBucketResult><Name>test-bucket</Name><KeyCount>0</KeyCount><MaxKeys>1000</MaxKeys><IsTruncated>false</IsTruncated></ListBucketResult>"#;
async fn observed_list_path(virtual_hosted: bool) -> String {
let server = wiremock::MockServer::start().await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.respond_with(
wiremock::ResponseTemplate::new(200)
.set_body_raw(EMPTY_LIST_XML, "application/xml"),
)
.mount(&server)
.await;
let storage = ObjectStorage::new(
&server.uri(),
"test-bucket",
"us-east-1",
Some("access"),
Some("secret"),
virtual_hosted,
);
assert!(!storage.health_check().await);
storage.refresh_total_size().await;
assert!(storage.health_check().await);
let requests = server.received_requests().await.unwrap();
requests[0].url.path().to_string()
}
#[tokio::test]
async fn test_path_style_addresses_bucket_in_path() {
assert_eq!(observed_list_path(false).await, "/test-bucket");
}
#[tokio::test]
async fn test_virtual_hosted_uses_endpoint_verbatim() {
assert_eq!(observed_list_path(true).await, "/");
}
#[test]
fn test_s3_total_size_returns_zero_before_init() {
let storage = ObjectStorage::new(
"http://localhost:9000",
"test-bucket",
"us-east-1",
Some("access"),
Some("secret"),
false,
);
assert!(!storage
.size_cache_initialized
.load(std::sync::atomic::Ordering::Relaxed));
}
#[test]
fn test_error_mapping_not_found() {
let err = object_store::Error::NotFound {
path: "test/key".to_string(),
source: "not found".into(),
};
match map_err(err) {
StorageError::NotFound => {}
other => panic!("Expected NotFound, got: {:?}", other),
}
}
#[test]
fn test_error_mapping_network() {
let err = object_store::Error::Generic {
store: "S3",
source: "connection refused".into(),
};
match map_err(err) {
StorageError::Network(msg) => {
assert!(msg.contains("connection refused"));
}
other => panic!("Expected Network, got: {:?}", other),
}
}
#[test]
fn test_encode_object_key() {
assert_eq!(encode_object_key("npm/@scope/pkg"), "npm/%40scope/pkg");
assert_eq!(
encode_object_key("npm/@babel/core/metadata.json"),
"npm/%40babel/core/metadata.json"
);
}
#[test]
fn test_decode_object_key_new_encoding() {
assert_eq!(decode_object_key("npm/%40scope/pkg"), "npm/@scope/pkg");
assert_eq!(
decode_object_key("npm/%40babel/core/metadata.json"),
"npm/@babel/core/metadata.json"
);
}
#[test]
fn test_decode_object_key_legacy_not_decoded() {
assert_eq!(decode_object_key("npm/_at_scope/pkg"), "npm/_at_scope/pkg");
assert_eq!(
decode_object_key("npm/_at_babel/core/metadata.json"),
"npm/_at_babel/core/metadata.json"
);
}
#[test]
fn test_encode_decode_roundtrip() {
let keys = [
"npm/@scope/pkg",
"npm/@babel/core/metadata.json",
"simple/key/no-at",
"raw/@org/file.txt",
"cargo/look_at_this/1.0.crate", "npm/some_at_pkg/metadata.json", ];
for key in keys {
assert_eq!(
decode_object_key(&encode_object_key(key)),
key,
"roundtrip failed for: {key}"
);
}
}
#[test]
fn test_no_roundtrip_collision_with_literal_at() {
let key = "cargo/look_at_this/1.0.crate";
let encoded = encode_object_key(key);
assert_eq!(encoded, key);
assert_eq!(decode_object_key(&encoded), key);
}
#[test]
fn test_encode_no_at() {
let key = "npm/chalk/metadata.json";
assert_eq!(encode_object_key(key), key);
}
#[test]
fn test_legacy_encode_for_fallback() {
assert_eq!(
encode_object_key_legacy("npm/@scope/pkg"),
"npm/_at_scope/pkg"
);
assert_eq!(
encode_object_key_legacy("npm/chalk/metadata.json"),
"npm/chalk/metadata.json"
);
}
}
#[cfg(test)]
mod proptests {
use super::*;
use proptest::prelude::*;
proptest! {
#[test]
fn object_key_roundtrip(key in "[a-z0-9@_./-]{1,100}") {
prop_assert_eq!(decode_object_key(&encode_object_key(&key)), key);
}
}
}