use bytes::Bytes;
use futures::Stream;
use serde::{Deserialize, Serialize};
use std::future::Future;
use std::path::{Component, Path, PathBuf};
use std::pin::Pin;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ObjectInfo {
pub key: String,
pub size: i64,
pub last_modified: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BucketInfo {
pub name: String,
pub region: Option<String>,
}
impl BucketInfo {
pub fn new(name: impl Into<String>, region: Option<String>) -> Self {
Self {
name: name.into(),
region,
}
}
}
#[derive(Debug, Clone)]
pub struct ListPage {
pub objects: Vec<ObjectInfo>,
pub truncated: bool,
}
#[non_exhaustive]
pub struct ObjectStream {
pub stream: Pin<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>>,
pub content_length: Option<i64>,
}
pub trait StorageBackend: Send + Sync {
fn list_buckets(
&self,
) -> Pin<Box<dyn Future<Output = Result<Vec<BucketInfo>, StorageError>> + Send + '_>>;
fn list_objects(
&self,
bucket: &str,
prefix: &str,
cap: usize,
) -> Pin<Box<dyn Future<Output = Result<ListPage, StorageError>> + Send + '_>>;
fn list_objects_all(
&self,
bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<ObjectInfo>, StorageError>> + Send + '_>>;
fn list_prefixes(
&self,
bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<String>, StorageError>> + Send + '_>>;
fn get_object(
&self,
bucket: &str,
key: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, StorageError>> + Send + '_>>;
fn put_object(
&self,
bucket: &str,
key: &str,
data: Vec<u8>,
) -> Pin<Box<dyn Future<Output = Result<(), StorageError>> + Send + '_>>;
fn get_object_stream(
&self,
bucket: &str,
key: &str,
) -> Pin<Box<dyn Future<Output = Result<ObjectStream, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let key = key.to_string();
Box::pin(async move {
let data = self.get_object(&bucket, &key).await?;
let content_length = Some(data.len() as i64);
let stream = futures::stream::once(async move { Ok(Bytes::from(data)) });
Ok(ObjectStream {
stream: Box::pin(stream),
content_length,
})
})
}
}
#[derive(Debug)]
pub enum StorageError {
NotFound(String),
Unauthorized,
AccountNotSignedUp,
WrongRegion,
BadRequest(String),
Other(String),
}
impl std::fmt::Display for StorageError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
StorageError::NotFound(msg) => write!(f, "not found: {msg}"),
StorageError::Unauthorized => {
write!(
f,
"credentials rejected by S3 (check keys, region, or expiry)"
)
}
StorageError::AccountNotSignedUp => {
write!(
f,
"the AWS account used for this request is not signed up for S3 — \
this usually means the request was signed with the wrong identity. \
Make sure you clicked Apply after pasting your credentials."
)
}
StorageError::WrongRegion => {
write!(
f,
"this bucket is in a different AWS region than the request was \
signed for. Set the region (or pick the bucket so its region is \
detected automatically) and try again."
)
}
StorageError::BadRequest(msg) => write!(f, "{msg}"),
StorageError::Other(msg) => write!(f, "{msg}"),
}
}
}
impl std::error::Error for StorageError {}
fn classify_s3_error<E, R>(err: &aws_sdk_s3::error::SdkError<E, R>) -> StorageError
where
E: std::error::Error + aws_sdk_s3::error::ProvideErrorMetadata + 'static,
R: std::fmt::Debug,
{
use aws_sdk_s3::error::ProvideErrorMetadata;
match err.code() {
Some(
"InvalidAccessKeyId"
| "SignatureDoesNotMatch"
| "ExpiredToken"
| "ExpiredTokenException"
| "InvalidToken"
| "AccessDenied"
| "AccessDeniedException"
| "UnrecognizedClientException"
| "InvalidClientTokenId"
| "AuthorizationHeaderMalformed",
) => StorageError::Unauthorized,
Some("NotSignedUp" | "OptInRequired") => StorageError::AccountNotSignedUp,
Some("PermanentRedirect" | "Redirect" | "IllegalLocationConstraintException") => {
StorageError::WrongRegion
}
Some("NoSuchBucket" | "NoSuchKey" | "NotFound") => {
StorageError::NotFound("the specified bucket or object does not exist".to_string())
}
Some("InvalidBucketName") => {
StorageError::BadRequest("the bucket name is not valid".to_string())
}
_ => {
tracing::warn!(
error = %aws_sdk_s3::error::DisplayErrorContext(err),
"unclassified S3 error"
);
StorageError::Other("could not complete the S3 request".to_string())
}
}
}
#[doc(hidden)]
#[derive(Clone)]
pub struct EphemeralS3Config {
pub http_client: aws_sdk_s3::config::SharedHttpClient,
pub endpoint_url: Option<String>,
pub force_path_style: bool,
}
const DEFAULT_REGION: &str = "us-east-1";
const OPERATION_ATTEMPT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
const OPERATION_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
pub struct S3Backend {
client: aws_sdk_s3::Client,
}
impl S3Backend {
pub async fn from_env() -> Self {
let config = aws_config::load_defaults(aws_config::BehaviorVersion::latest()).await;
Self::from_client(aws_sdk_s3::Client::new(&config))
}
pub async fn from_env_in_region(region: &str) -> Self {
let config = aws_config::defaults(aws_config::BehaviorVersion::latest())
.region(aws_sdk_s3::config::Region::new(region.to_string()))
.load()
.await;
Self::from_client(aws_sdk_s3::Client::new(&config))
}
pub fn from_client(client: aws_sdk_s3::Client) -> Self {
Self { client }
}
async fn fetch_bucket_details(&self) -> Result<Vec<BucketInfo>, StorageError> {
const MAX_BUCKETS: usize = 200;
let mut pages = self
.client
.list_buckets()
.max_buckets(MAX_BUCKETS as i32)
.into_paginator()
.send();
let mut buckets = Vec::new();
let mut truncated = false;
'pages: while let Some(page) = pages.next().await {
let page = page.map_err(|e| classify_s3_error(&e))?;
for bucket in page.buckets() {
if let Some(name) = bucket.name() {
buckets.push(BucketInfo::new(
name,
bucket.bucket_region().map(str::to_string),
));
}
if buckets.len() >= MAX_BUCKETS {
truncated = true;
break 'pages;
}
}
}
if truncated {
tracing::warn!(
max = MAX_BUCKETS,
"bucket listing truncated at cap; some buckets are not shown"
);
}
buckets.sort_by(|a, b| a.name.cmp(&b.name));
Ok(buckets)
}
pub fn from_credentials(
credentials: aws_sdk_s3::config::Credentials,
region: Option<&str>,
ephemeral: &Option<EphemeralS3Config>,
) -> Self {
Self::from_client(build_credentialed_client(credentials, region, ephemeral))
}
async fn list_objects_paginated(
&self,
bucket: &str,
prefix: &str,
cap: usize,
) -> Result<ListPage, StorageError> {
let mut pages = self
.client
.list_objects_v2()
.bucket(bucket)
.prefix(prefix)
.into_paginator()
.send();
let mut objects = Vec::new();
let mut truncated = false;
'pages: while let Some(page) = pages.next().await {
let page = page.map_err(|e| classify_s3_error(&e))?;
for obj in page.contents() {
if objects.len() >= cap {
truncated = true;
break 'pages;
}
if let Some(key) = obj.key() {
objects.push(ObjectInfo {
key: key.to_string(),
size: obj.size().unwrap_or(0),
last_modified: obj.last_modified().map(|t| t.to_string()),
});
}
}
}
if truncated {
tracing::warn!(
bucket = %bucket,
prefix = %prefix,
cap,
"object listing truncated at cap; some objects are not shown"
);
}
Ok(ListPage { objects, truncated })
}
}
pub fn build_credentialed_client(
credentials: aws_sdk_s3::config::Credentials,
region: Option<&str>,
ephemeral: &Option<EphemeralS3Config>,
) -> aws_sdk_s3::Client {
let region = region.unwrap_or(DEFAULT_REGION).to_string();
let timeouts = aws_sdk_s3::config::timeout::TimeoutConfig::builder()
.operation_attempt_timeout(OPERATION_ATTEMPT_TIMEOUT)
.operation_timeout(OPERATION_TIMEOUT)
.build();
let mut cfg = aws_sdk_s3::config::Builder::new()
.behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
.credentials_provider(credentials)
.timeout_config(timeouts)
.region(aws_sdk_s3::config::Region::new(region));
if let Some(e) = ephemeral {
cfg = cfg.http_client(e.http_client.clone());
if let Some(url) = &e.endpoint_url {
cfg = cfg.endpoint_url(url);
}
if e.force_path_style {
cfg = cfg.force_path_style(true);
}
}
aws_sdk_s3::Client::from_conf(cfg.build())
}
impl StorageBackend for S3Backend {
fn list_buckets(
&self,
) -> Pin<Box<dyn Future<Output = Result<Vec<BucketInfo>, StorageError>> + Send + '_>> {
Box::pin(self.fetch_bucket_details())
}
fn list_objects(
&self,
bucket: &str,
prefix: &str,
cap: usize,
) -> Pin<Box<dyn Future<Output = Result<ListPage, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let prefix = prefix.to_string();
Box::pin(async move { self.list_objects_paginated(&bucket, &prefix, cap).await })
}
fn list_objects_all(
&self,
bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<ObjectInfo>, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let prefix = prefix.to_string();
Box::pin(async move {
let mut objects = Vec::new();
let mut continuation: Option<String> = None;
loop {
let mut req = self
.client
.list_objects_v2()
.bucket(&bucket)
.prefix(&prefix);
if let Some(token) = continuation.take() {
req = req.continuation_token(token);
}
let resp = req.send().await.map_err(|e| classify_s3_error(&e))?;
for obj in resp.contents() {
if let Some(key) = obj.key() {
objects.push(ObjectInfo {
key: key.to_string(),
size: obj.size().unwrap_or(0),
last_modified: obj.last_modified().map(|t| t.to_string()),
});
}
}
if resp.is_truncated() == Some(true) {
match resp.next_continuation_token() {
Some(token) => continuation = Some(token.to_string()),
None => {
tracing::warn!(
bucket = %bucket,
prefix = %prefix,
returned = objects.len(),
"list_objects_all: response truncated but no continuation token; \
returning partial listing"
);
break;
}
}
} else {
break;
}
}
Ok(objects)
})
}
fn list_prefixes(
&self,
bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<String>, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let prefix = prefix.to_string();
Box::pin(async move {
const MAX_PREFIXES: usize = 1000;
let mut pages = self
.client
.list_objects_v2()
.bucket(&bucket)
.prefix(&prefix)
.delimiter("/")
.into_paginator()
.send();
let mut prefixes = Vec::new();
let mut truncated = false;
'pages: while let Some(page) = pages.next().await {
let page = page.map_err(|e| classify_s3_error(&e))?;
for cp in page.common_prefixes() {
if let Some(p) = cp.prefix() {
prefixes.push(p.to_string());
}
if prefixes.len() >= MAX_PREFIXES {
truncated = true;
break 'pages;
}
}
}
if truncated {
tracing::warn!(
bucket = %bucket,
prefix = %prefix,
max = MAX_PREFIXES,
"prefix listing truncated at cap; some child prefixes are not shown"
);
}
Ok(prefixes)
})
}
fn get_object(
&self,
bucket: &str,
key: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let key = key.to_string();
Box::pin(async move {
let resp = self
.client
.get_object()
.bucket(&bucket)
.key(&key)
.send()
.await
.map_err(|e| {
use aws_sdk_s3::operation::get_object::GetObjectError;
let classified = classify_s3_error(&e);
match e.into_service_error() {
GetObjectError::NoSuchKey(_) => {
StorageError::NotFound(format!("{bucket}/{key}"))
}
_ => classified,
}
})?;
let bytes = resp
.body
.collect()
.await
.map_err(|e| StorageError::Other(e.to_string()))?;
Ok(bytes.to_vec())
})
}
fn put_object(
&self,
bucket: &str,
key: &str,
data: Vec<u8>,
) -> Pin<Box<dyn Future<Output = Result<(), StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let key = key.to_string();
Box::pin(async move {
self.client
.put_object()
.bucket(&bucket)
.key(&key)
.body(aws_sdk_s3::primitives::ByteStream::from(data))
.send()
.await
.map_err(|e| {
use aws_sdk_s3::error::DisplayErrorContext;
StorageError::Other(format!("{}", DisplayErrorContext(&e)))
})?;
Ok(())
})
}
fn get_object_stream(
&self,
bucket: &str,
key: &str,
) -> Pin<Box<dyn Future<Output = Result<ObjectStream, StorageError>> + Send + '_>> {
let bucket = bucket.to_string();
let key = key.to_string();
Box::pin(async move {
let resp = self
.client
.get_object()
.bucket(&bucket)
.key(&key)
.send()
.await
.map_err(|e| {
use aws_sdk_s3::operation::get_object::GetObjectError;
let classified = classify_s3_error(&e);
match e.into_service_error() {
GetObjectError::NoSuchKey(_) => {
StorageError::NotFound(format!("{bucket}/{key}"))
}
_ => classified,
}
})?;
let content_length = resp.content_length();
let stream = futures::stream::unfold(resp.body, |mut body| async move {
match body.next().await {
Some(Ok(chunk)) => Some((Ok(chunk), body)),
Some(Err(e)) => Some((Err(std::io::Error::other(e)), body)),
None => None,
}
});
Ok(ObjectStream {
stream: Box::pin(stream),
content_length,
})
})
}
}
#[cfg(unix)]
fn make_temporary_root_private(root: &Path) -> std::io::Result<()> {
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(root, std::fs::Permissions::from_mode(0o700))
}
#[cfg(not(unix))]
fn make_temporary_root_private(_root: &Path) -> std::io::Result<()> {
Ok(())
}
#[cfg(not(windows))]
fn replace_file_atomically(source: &Path, destination: &Path) -> std::io::Result<()> {
std::fs::rename(source, destination)
}
#[cfg(windows)]
fn replace_file_atomically(source: &Path, destination: &Path) -> std::io::Result<()> {
use std::os::windows::ffi::OsStrExt;
use windows_sys::Win32::Storage::FileSystem::{
MOVEFILE_REPLACE_EXISTING, MOVEFILE_WRITE_THROUGH, MoveFileExW,
};
let source: Vec<u16> = source.as_os_str().encode_wide().chain([0]).collect();
let destination: Vec<u16> = destination.as_os_str().encode_wide().chain([0]).collect();
let result = unsafe {
MoveFileExW(
source.as_ptr(),
destination.as_ptr(),
MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH,
)
};
if result == 0 {
Err(std::io::Error::last_os_error())
} else {
Ok(())
}
}
pub struct LocalBackend {
root: PathBuf,
remove_on_drop: bool,
}
impl LocalBackend {
pub fn new(root: impl Into<PathBuf>) -> Self {
let root = root.into();
let root = root.canonicalize().unwrap_or(root);
Self {
root,
remove_on_drop: false,
}
}
pub(crate) fn new_temporary_aggregate() -> Self {
Self::new_temporary_aggregate_in(std::env::temp_dir())
}
fn new_temporary_aggregate_in(base: impl AsRef<Path>) -> Self {
let base = base.as_ref();
let base = base.canonicalize().unwrap_or_else(|_| base.to_path_buf());
let root = base.join(format!(
"dial9-aggregate-{}",
uuid::Uuid::new_v4().as_hyphenated()
));
Self {
root,
remove_on_drop: true,
}
}
pub(crate) fn root(&self) -> &Path {
&self.root
}
}
impl Drop for LocalBackend {
fn drop(&mut self) {
if !self.remove_on_drop {
return;
}
match std::fs::remove_dir_all(&self.root) {
Ok(()) => {}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {}
Err(e) => {
tracing::warn!(
path = %self.root.display(),
error = %e,
"failed to remove temporary aggregate directory"
);
}
}
}
}
impl StorageBackend for LocalBackend {
fn list_buckets(
&self,
) -> Pin<Box<dyn Future<Output = Result<Vec<BucketInfo>, StorageError>> + Send + '_>> {
Box::pin(async { Ok(Vec::new()) })
}
fn list_objects(
&self,
_bucket: &str,
prefix: &str,
cap: usize,
) -> Pin<Box<dyn Future<Output = Result<ListPage, StorageError>> + Send + '_>> {
let prefix = prefix.to_string();
Box::pin(async move {
let root = self.root.clone();
let prefix2 = prefix.clone();
tokio::task::spawn_blocking(move || {
let mut objects = Vec::new();
collect_files(&root, &root, &prefix2, &mut objects, 0, &mut 0)?;
objects.sort_by(|a, b| a.key.cmp(&b.key));
let truncated = objects.len() > cap;
objects.truncate(cap);
Ok(ListPage { objects, truncated })
})
.await
.map_err(|e| StorageError::Other(e.to_string()))?
})
}
fn list_objects_all(
&self,
_bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<ObjectInfo>, StorageError>> + Send + '_>> {
let prefix = prefix.to_string();
Box::pin(async move {
let root = self.root.clone();
tokio::task::spawn_blocking(move || {
let mut objects = Vec::new();
collect_files_uncapped(&root, &root, &prefix, &mut objects)?;
objects.sort_by(|a, b| a.key.cmp(&b.key));
Ok(objects)
})
.await
.map_err(|e| StorageError::Other(e.to_string()))?
})
}
fn list_prefixes(
&self,
_bucket: &str,
prefix: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<String>, StorageError>> + Send + '_>> {
let prefix = prefix.to_string();
Box::pin(async move {
let root = self.root.clone();
let prefix2 = prefix.clone();
tokio::task::spawn_blocking(move || {
let dir = root.join(&prefix2);
let dir = match dir.canonicalize() {
Ok(d) if d.starts_with(&root) => d,
Ok(_) => {
return Err(StorageError::NotFound(
"path escapes root directory".to_string(),
));
}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(vec![]),
Err(e) => return Err(StorageError::Other(e.to_string())),
};
let entries = match std::fs::read_dir(&dir) {
Ok(e) => e,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(vec![]),
Err(e) => return Err(StorageError::Other(e.to_string())),
};
let mut prefixes = Vec::new();
for entry in entries {
let entry = entry.map_err(|e| StorageError::Other(e.to_string()))?;
let path = entry.path();
let canonical = match path.canonicalize() {
Ok(c) if c.starts_with(&root) => c,
_ => continue,
};
if canonical.is_dir() {
let name = entry.file_name().to_string_lossy().into_owned();
prefixes.push(format!("{prefix2}{name}/"));
}
}
prefixes.sort();
Ok(prefixes)
})
.await
.map_err(|e| StorageError::Other(e.to_string()))?
})
}
fn get_object(
&self,
_bucket: &str,
key: &str,
) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, StorageError>> + Send + '_>> {
let path = self.root.join(key);
let root = self.root.clone();
Box::pin(async move {
tokio::task::spawn_blocking(move || {
let canonical = path.canonicalize().map_err(|e| match e.kind() {
std::io::ErrorKind::NotFound => {
StorageError::NotFound(path.display().to_string())
}
_ => StorageError::Other(e.to_string()),
})?;
if !canonical.starts_with(&root) {
return Err(StorageError::NotFound(
"path escapes root directory".to_string(),
));
}
std::fs::read(&canonical).map_err(|e| match e.kind() {
std::io::ErrorKind::NotFound => {
StorageError::NotFound(path.display().to_string())
}
_ => StorageError::Other(e.to_string()),
})
})
.await
.map_err(|e| StorageError::Other(e.to_string()))?
})
}
fn put_object(
&self,
_bucket: &str,
key: &str,
data: Vec<u8>,
) -> Pin<Box<dyn Future<Output = Result<(), StorageError>> + Send + '_>> {
let root = self.root.clone();
let private_temporary_root = self.remove_on_drop;
let key = key.to_string();
Box::pin(async move {
tokio::task::spawn_blocking(move || {
let path = root.join(&key);
if path.components().any(|c| c == Component::ParentDir) {
return Err(StorageError::Other(
"key contains path traversal".to_string(),
));
}
let parent = path.parent().ok_or_else(|| {
StorageError::Other("object path has no parent directory".to_string())
})?;
std::fs::create_dir_all(parent).map_err(|e| StorageError::Other(e.to_string()))?;
if private_temporary_root {
make_temporary_root_private(&root)
.map_err(|e| StorageError::Other(e.to_string()))?;
}
let canonical_parent = parent
.canonicalize()
.map_err(|e| StorageError::Other(e.to_string()))?;
if !canonical_parent.starts_with(&root) {
return Err(StorageError::Other(
"path escapes root directory".to_string(),
));
}
let temp_path = canonical_parent.join(format!(
".dial9-write-{}",
uuid::Uuid::new_v4().as_hyphenated()
));
std::fs::write(&temp_path, data).map_err(|e| StorageError::Other(e.to_string()))?;
if let Err(rename_error) = replace_file_atomically(&temp_path, &path) {
if let Err(cleanup_error) = std::fs::remove_file(&temp_path)
&& cleanup_error.kind() != std::io::ErrorKind::NotFound
{
tracing::warn!(
path = %temp_path.display(),
error = %cleanup_error,
"failed to clean up temporary object write"
);
}
return Err(StorageError::Other(rename_error.to_string()));
}
Ok(())
})
.await
.map_err(|e| StorageError::Other(e.to_string()))?
})
}
}
const MAX_COLLECT_DEPTH: u32 = 10;
const MAX_COLLECT_FILES: usize = 50;
const MAX_ENTRIES_VISITED: usize = 500;
fn is_skipped_dir(name: &str) -> bool {
name.starts_with('.') || matches!(name, "target" | "node_modules")
}
fn collect_files(
root: &Path,
dir: &Path,
prefix: &str,
out: &mut Vec<ObjectInfo>,
depth: u32,
visited: &mut usize,
) -> Result<(), StorageError> {
if depth > MAX_COLLECT_DEPTH
|| out.len() >= MAX_COLLECT_FILES
|| *visited >= MAX_ENTRIES_VISITED
{
return Ok(());
}
let entries = match std::fs::read_dir(dir) {
Ok(e) => e,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::PermissionDenied => {
return Err(StorageError::Other("permission denied".into()));
}
Err(e) => return Err(StorageError::Other(e.to_string())),
};
for entry in entries {
*visited += 1;
if out.len() >= MAX_COLLECT_FILES || *visited >= MAX_ENTRIES_VISITED {
break;
}
let entry = entry.map_err(|e| StorageError::Other(e.to_string()))?;
let path = entry.path();
let canonical = match path.canonicalize() {
Ok(c) if c.starts_with(root) => c,
_ => continue,
};
if canonical.is_dir() {
let name = entry.file_name();
let name = name.to_string_lossy();
if !is_skipped_dir(&name) {
collect_files(root, &canonical, prefix, out, depth + 1, visited)?;
}
} else if canonical.is_file() {
let file_name = entry.file_name();
let file_name_str = file_name.to_string_lossy();
if file_name_str.starts_with('.') {
continue;
}
let key = path
.strip_prefix(root)
.unwrap_or(&path)
.to_string_lossy()
.into_owned();
if key.starts_with(prefix) {
let meta = std::fs::metadata(&canonical)
.map_err(|e| StorageError::Other(e.to_string()))?;
out.push(ObjectInfo {
key,
size: meta.len() as i64,
last_modified: meta.modified().ok().and_then(|t| {
t.duration_since(std::time::UNIX_EPOCH)
.ok()
.map(|d| d.as_secs().to_string())
}),
});
}
}
}
Ok(())
}
fn collect_files_uncapped(
root: &Path,
dir: &Path,
prefix: &str,
out: &mut Vec<ObjectInfo>,
) -> Result<(), StorageError> {
let entries = match std::fs::read_dir(dir) {
Ok(e) => e,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::PermissionDenied => {
return Err(StorageError::Other("permission denied".into()));
}
Err(e) => return Err(StorageError::Other(e.to_string())),
};
for entry in entries {
let entry = entry.map_err(|e| StorageError::Other(e.to_string()))?;
let path = entry.path();
let canonical = match path.canonicalize() {
Ok(c) if c.starts_with(root) => c,
_ => continue,
};
if canonical.is_dir() {
let name = entry.file_name();
let name = name.to_string_lossy();
if !is_skipped_dir(&name) {
collect_files_uncapped(root, &canonical, prefix, out)?;
}
} else if canonical.is_file() {
let key = path
.strip_prefix(root)
.unwrap_or(&path)
.to_string_lossy()
.into_owned();
if key.starts_with(prefix) {
let meta = std::fs::metadata(&canonical)
.map_err(|e| StorageError::Other(e.to_string()))?;
out.push(ObjectInfo {
key,
size: meta.len() as i64,
last_modified: meta.modified().ok().and_then(|t| {
t.duration_since(std::time::UNIX_EPOCH)
.ok()
.map(|d| d.as_secs().to_string())
}),
});
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn replay_backend(
responses: Vec<&str>,
) -> (
S3Backend,
aws_smithy_http_client::test_util::StaticReplayClient,
) {
use aws_smithy_http_client::test_util::{ReplayEvent, StaticReplayClient};
use aws_smithy_types::body::SdkBody;
let events = responses
.into_iter()
.map(|body| {
ReplayEvent::new(
http::Request::builder()
.uri("https://s3.amazonaws.com/")
.body(SdkBody::empty())
.unwrap(),
http::Response::builder()
.status(200)
.body(SdkBody::from(body))
.unwrap(),
)
})
.collect();
let http_client = StaticReplayClient::new(events);
let cfg = aws_sdk_s3::config::Builder::new()
.behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
.credentials_provider(aws_sdk_s3::config::Credentials::new(
"test", "test", None, None, "test",
))
.region(aws_sdk_s3::config::Region::new("us-east-1"))
.http_client(http_client.clone())
.build();
(
S3Backend::from_client(aws_sdk_s3::Client::from_conf(cfg)),
http_client,
)
}
fn replay_error_backend(status: u16, body: &str) -> S3Backend {
use aws_smithy_http_client::test_util::{ReplayEvent, StaticReplayClient};
use aws_smithy_types::body::SdkBody;
let http_client = StaticReplayClient::new(vec![ReplayEvent::new(
http::Request::builder()
.uri("https://s3.amazonaws.com/")
.body(SdkBody::empty())
.unwrap(),
http::Response::builder()
.status(status)
.body(SdkBody::from(body))
.unwrap(),
)]);
let cfg = aws_sdk_s3::config::Builder::new()
.behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
.credentials_provider(aws_sdk_s3::config::Credentials::new(
"test", "test", None, None, "test",
))
.region(aws_sdk_s3::config::Region::new("us-east-1"))
.http_client(http_client)
.build();
S3Backend::from_client(aws_sdk_s3::Client::from_conf(cfg))
}
#[tokio::test]
async fn list_buckets_requests_and_preserves_regions() {
use aws_smithy_http_client::test_util::infallible_client_fn;
let http_client = infallible_client_fn(|req: http::Request<_>| {
assert!(
req.uri()
.query()
.is_some_and(|query| query.contains("max-buckets=200")),
"ListBuckets must include a valid parameter so S3 returns BucketRegion: {}",
req.uri()
);
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<ListAllMyBucketsResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/">
<Buckets>
<Bucket>
<Name>dial9-cape-town</Name>
<CreationDate>2026-07-21T00:00:00.000Z</CreationDate>
<BucketRegion>af-south-1</BucketRegion>
</Bucket>
<Bucket>
<Name>dial9-oregon</Name>
<CreationDate>2026-07-21T00:00:00.000Z</CreationDate>
<BucketRegion>us-west-2</BucketRegion>
</Bucket>
</Buckets>
</ListAllMyBucketsResult>"#;
http::Response::builder()
.status(200)
.header("content-type", "application/xml")
.body(body)
.unwrap()
});
let cfg = aws_sdk_s3::config::Builder::new()
.behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
.credentials_provider(aws_sdk_s3::config::Credentials::new(
"test", "test", None, None, "test",
))
.region(aws_sdk_s3::config::Region::new("us-east-1"))
.http_client(http_client)
.build();
let backend = S3Backend::from_client(aws_sdk_s3::Client::from_conf(cfg));
let buckets = backend.list_buckets().await.unwrap();
assert_eq!(
buckets,
vec![
BucketInfo::new("dial9-cape-town", Some("af-south-1".to_string())),
BucketInfo::new("dial9-oregon", Some("us-west-2".to_string())),
]
);
}
#[tokio::test]
async fn permanent_redirect_classifies_as_wrong_region() {
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<Error>
<Code>PermanentRedirect</Code>
<Message>The bucket you are attempting to access must be addressed using the specified endpoint.</Message>
<Endpoint>my-bucket.s3.us-west-2.amazonaws.com</Endpoint>
</Error>"#;
let backend = replay_error_backend(301, body);
let err = backend
.list_prefixes("my-bucket", "")
.await
.expect_err("a PermanentRedirect must surface as an error");
assert!(
matches!(err, StorageError::WrongRegion),
"expected WrongRegion, got {err:?}"
);
assert!(err.to_string().contains("region"), "message: {err}");
}
#[tokio::test]
async fn illegal_location_constraint_classifies_as_wrong_region() {
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<Error>
<Code>IllegalLocationConstraintException</Code>
<Message>The af-south-1 location constraint is incompatible for the region specific endpoint this request was sent to.</Message>
</Error>"#;
let backend = replay_error_backend(400, body);
let err = backend
.list_prefixes("my-bucket", "")
.await
.expect_err("an IllegalLocationConstraintException must surface as an error");
assert!(
matches!(err, StorageError::WrongRegion),
"expected WrongRegion, got {err:?}"
);
assert!(err.to_string().contains("region"), "message: {err}");
}
#[tokio::test]
async fn no_such_bucket_classifies_as_not_found() {
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<Error>
<Code>NoSuchBucket</Code>
<Message>The specified bucket does not exist</Message>
<BucketName>test-shanks</BucketName>
</Error>"#;
let backend = replay_error_backend(404, body);
let err = backend
.list_prefixes("test-shanks", "")
.await
.expect_err("a NoSuchBucket must surface as an error");
assert!(
matches!(err, StorageError::NotFound(_)),
"expected NotFound, got {err:?}"
);
}
#[tokio::test]
async fn invalid_bucket_name_classifies_as_bad_request() {
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<Error>
<Code>InvalidBucketName</Code>
<Message>The specified bucket is not valid.</Message>
<BucketName>Not_A_Valid_Bucket</BucketName>
</Error>"#;
let backend = replay_error_backend(400, body);
let err = backend
.list_prefixes("Not_A_Valid_Bucket", "")
.await
.expect_err("an InvalidBucketName must surface as an error");
assert!(
matches!(err, StorageError::BadRequest(_)),
"expected BadRequest, got {err:?}"
);
}
#[tokio::test]
async fn list_prefixes_follows_continuation_token() {
let page1 = r#"<?xml version="1.0" encoding="UTF-8"?>
<ListBucketResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/">
<Name>bucket</Name><Prefix></Prefix><Delimiter>/</Delimiter>
<IsTruncated>true</IsTruncated>
<NextContinuationToken>TOKEN_A</NextContinuationToken>
<CommonPrefixes><Prefix>a/</Prefix></CommonPrefixes>
<CommonPrefixes><Prefix>b/</Prefix></CommonPrefixes>
</ListBucketResult>"#;
let page2 = r#"<?xml version="1.0" encoding="UTF-8"?>
<ListBucketResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/">
<Name>bucket</Name><Prefix></Prefix><Delimiter>/</Delimiter>
<IsTruncated>false</IsTruncated>
<CommonPrefixes><Prefix>c/</Prefix></CommonPrefixes>
</ListBucketResult>"#;
let (backend, http_client) = replay_backend(vec![page1, page2]);
let prefixes = backend.list_prefixes("bucket", "").await.unwrap();
assert_eq!(prefixes, vec!["a/", "b/", "c/"]);
let requests = http_client.actual_requests().collect::<Vec<_>>();
assert_eq!(requests.len(), 2, "expected two list calls");
assert!(
requests[1].uri().contains("continuation-token=TOKEN_A"),
"second request must carry the continuation token, got: {}",
requests[1].uri()
);
}
#[tokio::test]
async fn put_object_rejects_path_traversal_without_creating_dirs() {
let outer = tempfile::tempdir().unwrap();
let root = outer.path().join("root");
std::fs::create_dir(&root).unwrap();
let backend = LocalBackend::new(&root);
let err = backend
.put_object("bucket", "../escape/evil.bin", b"x".to_vec())
.await
.expect_err("traversal key must be rejected");
match err {
StorageError::Other(msg) => {
assert!(msg.contains("path traversal"), "unexpected message: {msg}")
}
other => panic!("expected StorageError::Other, got {other:?}"),
}
let escape_dir = outer.path().join("escape");
assert!(
!escape_dir.exists(),
"path traversal created a directory outside the root: {}",
escape_dir.display()
);
}
#[tokio::test]
async fn put_object_writes_normal_key() {
let dir = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(dir.path());
backend
.put_object("bucket", "a/b/c.bin", b"hi".to_vec())
.await
.unwrap();
let written = dir.path().join("a/b/c.bin");
assert!(written.exists(), "expected file at {}", written.display());
assert_eq!(std::fs::read(&written).unwrap(), b"hi");
}
#[tokio::test]
async fn concurrent_put_object_readers_see_only_complete_versions() {
let dir = tempfile::tempdir().unwrap();
let backend = std::sync::Arc::new(LocalBackend::new(dir.path()));
let key = "samples/part.parquet";
let body_len = 1024 * 1024;
backend
.put_object("local", key, vec![0; body_len])
.await
.unwrap();
let mut writers = Vec::new();
for byte in 1..=8u8 {
let backend = std::sync::Arc::clone(&backend);
writers.push(tokio::spawn(async move {
backend
.put_object("local", key, vec![byte; body_len])
.await
.unwrap();
}));
}
for _ in 0..64 {
let body = backend.get_object("local", key).await.unwrap();
assert_eq!(body.len(), body_len, "reader observed a truncated object");
assert!(
body.iter().all(|byte| *byte == body[0]),
"reader observed bytes from multiple object versions"
);
tokio::task::yield_now().await;
}
for writer in writers {
writer.await.unwrap();
}
}
#[tokio::test]
async fn temporary_aggregate_backend_cleans_up_on_drop() {
let backend = LocalBackend::new_temporary_aggregate();
let root = backend.root.clone();
backend
.put_object("cache", "a/b/c.parquet", b"cached".to_vec())
.await
.unwrap();
assert!(root.join("a/b/c.parquet").exists());
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mode = std::fs::metadata(&root).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o700, "temporary cache root must be owner-only");
}
drop(backend);
assert!(!root.exists(), "temporary aggregate directory leaked");
}
#[cfg(unix)]
#[tokio::test]
async fn temporary_aggregate_backend_resolves_symlinked_temp_base() {
use std::os::unix::fs::symlink;
let outer = tempfile::tempdir().unwrap();
let real_base = outer.path().join("real-temp");
let linked_base = outer.path().join("temp-link");
std::fs::create_dir(&real_base).unwrap();
symlink(&real_base, &linked_base).unwrap();
let backend = LocalBackend::new_temporary_aggregate_in(&linked_base);
assert!(backend.root.starts_with(real_base.canonicalize().unwrap()));
backend
.put_object("cache", "samples/part.parquet", b"complete".to_vec())
.await
.unwrap();
assert_eq!(
backend
.get_object("cache", "samples/part.parquet")
.await
.unwrap(),
b"complete"
);
}
fn replay_backend_status(status: u16, body: &str) -> S3Backend {
use aws_smithy_http_client::test_util::{ReplayEvent, StaticReplayClient};
use aws_smithy_types::body::SdkBody;
let events = vec![ReplayEvent::new(
http::Request::builder()
.uri("https://s3.amazonaws.com/")
.body(SdkBody::empty())
.unwrap(),
http::Response::builder()
.status(status)
.body(SdkBody::from(body))
.unwrap(),
)];
let http_client = StaticReplayClient::new(events);
let cfg = aws_sdk_s3::config::Builder::new()
.behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
.credentials_provider(aws_sdk_s3::config::Credentials::new(
"test", "test", None, None, "test",
))
.region(aws_sdk_s3::config::Region::new("us-east-1"))
.http_client(http_client)
.build();
S3Backend::from_client(aws_sdk_s3::Client::from_conf(cfg))
}
#[tokio::test]
async fn list_objects_all_maps_auth_failure_to_unauthorized() {
let body = r#"<?xml version="1.0" encoding="UTF-8"?>
<Error><Code>InvalidAccessKeyId</Code>
<Message>The AWS Access Key Id you provided does not exist in our records.</Message>
</Error>"#;
let backend = replay_backend_status(403, body);
let err = backend
.list_objects_all("bucket", "prefix")
.await
.expect_err("auth failure must be an error");
assert!(
matches!(err, StorageError::Unauthorized),
"expected Unauthorized, got {err:?}"
);
}
#[test]
fn collect_files_caps_entries_visited() {
let dir = tempfile::tempdir().unwrap();
let n = MAX_ENTRIES_VISITED + 500;
for i in 0..n {
std::fs::write(dir.path().join(format!("file_{i:05}.bin")), b"x").unwrap();
}
let mut out = Vec::new();
let mut visited = 0;
collect_files(dir.path(), dir.path(), "", &mut out, 0, &mut visited).unwrap();
assert!(
visited <= MAX_ENTRIES_VISITED,
"visited {visited} entries, expected at most {MAX_ENTRIES_VISITED}"
);
assert!(
out.len() <= MAX_COLLECT_FILES,
"collected {} files, expected at most {MAX_COLLECT_FILES}",
out.len()
);
}
}