use async_trait::async_trait;
use eyre::{Context, Result};
use std::path::{Path, PathBuf};
#[async_trait]
pub trait UserStorageBackend: Send + Sync {
fn backend_name(&self) -> &'static str {
"custom"
}
async fn put(&self, key: &str, data: &[u8]) -> Result<()>;
async fn get(&self, key: &str) -> Result<Vec<u8>>;
async fn get_optional(&self, key: &str) -> Result<Option<Vec<u8>>>;
async fn exists(&self, key: &str) -> Result<bool>;
async fn delete(&self, key: &str) -> Result<()>;
async fn list(&self, prefix: &str) -> Result<Vec<String>>;
}
#[derive(Debug, Clone)]
pub struct LocalBackend {
base_dir: PathBuf,
}
impl LocalBackend {
pub fn new(base_dir: impl AsRef<Path>) -> Self {
Self {
base_dir: base_dir.as_ref().to_path_buf(),
}
}
fn validate_key(key: &str) -> Result<()> {
if key.is_empty() {
return Err(eyre::eyre!("storage key cannot be empty"));
}
let path = Path::new(key);
if path.is_absolute() {
return Err(eyre::eyre!("storage key cannot be absolute: {key}"));
}
for component in path.components() {
use std::path::Component;
match component {
Component::CurDir | Component::ParentDir => {
return Err(eyre::eyre!("storage key contains path traversal: {key}"));
}
Component::Prefix(_) | Component::RootDir => {
return Err(eyre::eyre!("storage key contains root/prefix component: {key}"));
}
Component::Normal(_) => {}
}
}
Ok(())
}
fn resolve(&self, key: &str) -> Result<PathBuf> {
Self::validate_key(key)?;
let joined = self.base_dir.join(key);
if !joined.starts_with(&self.base_dir) {
return Err(eyre::eyre!("storage key escapes base_dir: {key}"));
}
Ok(joined)
}
}
#[async_trait]
impl UserStorageBackend for LocalBackend {
fn backend_name(&self) -> &'static str {
"local"
}
async fn put(&self, key: &str, data: &[u8]) -> Result<()> {
let path = self.resolve(key).wrap_err_with(|| format!("resolve key {key}"))?;
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent).await.wrap_err("create parent dirs")?;
}
tokio::fs::write(&path, data)
.await
.wrap_err_with(|| format!("write {}", path.display()))
}
async fn get(&self, key: &str) -> Result<Vec<u8>> {
let path = self.resolve(key).wrap_err_with(|| format!("resolve key {key}"))?;
tokio::fs::read(&path)
.await
.wrap_err_with(|| format!("read {}", path.display()))
}
async fn get_optional(&self, key: &str) -> Result<Option<Vec<u8>>> {
let path = self.resolve(key).wrap_err_with(|| format!("resolve key {key}"))?;
match tokio::fs::read(&path).await {
Ok(data) => Ok(Some(data)),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(e) => Err(e).wrap_err_with(|| format!("read {}", path.display())),
}
}
async fn exists(&self, key: &str) -> Result<bool> {
let path = self.resolve(key).wrap_err_with(|| format!("resolve key {key}"))?;
Ok(tokio::fs::try_exists(&path)
.await
.wrap_err_with(|| format!("stat {}", path.display()))?)
}
async fn delete(&self, key: &str) -> Result<()> {
let path = self.resolve(key).wrap_err_with(|| format!("resolve key {key}"))?;
match tokio::fs::remove_file(&path).await {
Ok(()) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(e) => Err(e).wrap_err_with(|| format!("delete {}", path.display())),
}
}
async fn list(&self, prefix: &str) -> Result<Vec<String>> {
let root = self
.resolve(prefix)
.wrap_err_with(|| format!("resolve prefix {prefix}"))?;
let mut out = Vec::new();
let mut stack = vec![root];
while let Some(dir) = stack.pop() {
let mut rd = match tokio::fs::read_dir(&dir).await {
Ok(rd) => rd,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => continue,
Err(e) => return Err(e).wrap_err_with(|| format!("read_dir {}", dir.display())),
};
while let Some(entry) = rd.next_entry().await.wrap_err("read dir entry")? {
let path = entry.path();
if entry.file_type().await.wrap_err("stat entry")?.is_dir() {
stack.push(path);
} else if let Ok(rel) = path.strip_prefix(&self.base_dir) {
out.push(rel.to_string_lossy().replace('\\', "/"));
}
}
}
Ok(out)
}
}
#[cfg(feature = "persisted-immutable-data-s3")]
#[derive(Debug, Clone)]
pub struct S3Backend {
client: aws_sdk_s3::Client,
bucket: String,
}
#[cfg(feature = "persisted-immutable-data-s3")]
fn is_s3_object_not_found(error_code: Option<&str>, status: Option<u16>) -> bool {
match error_code {
Some("NoSuchKey" | "NotFound") => true,
Some(_) => false,
None => status == Some(404),
}
}
#[cfg(feature = "persisted-immutable-data-s3")]
impl S3Backend {
pub async fn new(bucket: String, region: String) -> Result<Self> {
let config = aws_config::defaults(aws_config::BehaviorVersion::latest())
.region(aws_config::Region::new(region))
.load()
.await;
let client = aws_sdk_s3::Client::new(&config);
Ok(Self { client, bucket })
}
pub fn from_sdk_config(config: &aws_types::SdkConfig, bucket: String) -> Self {
Self {
client: aws_sdk_s3::Client::new(config),
bucket,
}
}
}
#[cfg(feature = "persisted-immutable-data-s3")]
#[async_trait]
impl UserStorageBackend for S3Backend {
fn backend_name(&self) -> &'static str {
"s3"
}
async fn put(&self, key: &str, data: &[u8]) -> Result<()> {
use aws_sdk_s3::primitives::ByteStream;
self.client
.put_object()
.bucket(&self.bucket)
.key(key)
.body(ByteStream::from(data.to_vec()))
.send()
.await
.wrap_err_with(|| format!("s3 put s3://{}/{}", self.bucket, key))?;
Ok(())
}
async fn get(&self, key: &str) -> Result<Vec<u8>> {
let resp = self
.client
.get_object()
.bucket(&self.bucket)
.key(key)
.send()
.await
.wrap_err_with(|| format!("s3 get s3://{}/{}", self.bucket, key))?;
let data = resp
.body
.collect()
.await
.wrap_err("read s3 response body")?
.into_bytes()
.to_vec();
Ok(data)
}
async fn get_optional(&self, key: &str) -> Result<Option<Vec<u8>>> {
let response = match self.client.get_object().bucket(&self.bucket).key(key).send().await {
Ok(response) => response,
Err(error) => {
use aws_sdk_s3::error::ProvideErrorMetadata;
let error_code = error.as_service_error().and_then(ProvideErrorMetadata::code);
let status = error.raw_response().map(|response| response.status().as_u16());
let is_not_found = is_s3_object_not_found(error_code, status);
if is_not_found {
return Ok(None);
}
return Err(error).wrap_err_with(|| format!("s3 get s3://{}/{}", self.bucket, key));
}
};
let data = response
.body
.collect()
.await
.wrap_err("read s3 response body")?
.into_bytes()
.to_vec();
Ok(Some(data))
}
async fn exists(&self, key: &str) -> Result<bool> {
match self.client.head_object().bucket(&self.bucket).key(key).send().await {
Ok(_) => Ok(true),
Err(e) => {
use aws_sdk_s3::error::ProvideErrorMetadata;
let error_code = e.as_service_error().and_then(ProvideErrorMetadata::code);
let status = e.raw_response().map(|response| response.status().as_u16());
let is_not_found = is_s3_object_not_found(error_code, status);
if is_not_found {
Ok(false)
} else {
Err(e).wrap_err_with(|| format!("s3 head s3://{}/{}", self.bucket, key))
}
}
}
}
async fn delete(&self, key: &str) -> Result<()> {
self.client
.delete_object()
.bucket(&self.bucket)
.key(key)
.send()
.await
.wrap_err_with(|| format!("s3 delete s3://{}/{}", self.bucket, key))?;
Ok(())
}
async fn list(&self, prefix: &str) -> Result<Vec<String>> {
let mut out = Vec::new();
let mut continuation: Option<String> = None;
loop {
let mut req = self.client.list_objects_v2().bucket(&self.bucket).prefix(prefix);
if let Some(token) = &continuation {
req = req.continuation_token(token);
}
let resp = req
.send()
.await
.wrap_err_with(|| format!("s3 list s3://{}/{}", self.bucket, prefix))?;
for obj in resp.contents() {
if let Some(key) = obj.key() {
out.push(key.to_string());
}
}
match resp.next_continuation_token() {
Some(token) => continuation = Some(token.to_string()),
None => break,
}
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "persisted-immutable-data-s3")]
#[test]
fn s3_bucket_errors_are_not_classified_as_object_misses() {
assert!(is_s3_object_not_found(Some("NoSuchKey"), Some(404)));
assert!(is_s3_object_not_found(Some("NotFound"), Some(404)));
assert!(is_s3_object_not_found(None, Some(404)));
assert!(!is_s3_object_not_found(Some("NoSuchBucket"), Some(404)));
assert!(!is_s3_object_not_found(Some("AccessDenied"), Some(404)));
assert!(!is_s3_object_not_found(None, Some(500)));
}
#[tokio::test]
async fn local_backend_roundtrip() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let key = "test/data.bin";
let data = b"hello pid";
backend.put(key, data).await.unwrap();
assert!(backend.exists(key).await.unwrap());
let downloaded = backend.get(key).await.unwrap();
assert_eq!(downloaded, data);
}
#[tokio::test]
async fn local_backend_missing() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
assert!(!backend.exists("missing").await.unwrap());
assert!(backend.get("missing").await.is_err());
assert!(backend.get_optional("missing").await.unwrap().is_none());
}
#[tokio::test]
async fn local_backend_delete_idempotent() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
backend.put("obj/a", b"x").await.unwrap();
assert!(backend.exists("obj/a").await.unwrap());
backend.delete("obj/a").await.unwrap();
assert!(!backend.exists("obj/a").await.unwrap());
backend.delete("obj/a").await.unwrap();
}
#[tokio::test]
async fn local_backend_list_prefix() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
backend.put("owner/1/a.wasm", b"a").await.unwrap();
backend.put("owner/1/b.rego", b"b").await.unwrap();
backend.put("owner/2/c.wasm", b"c").await.unwrap();
let mut listed = backend.list("owner/1").await.unwrap();
listed.sort();
assert_eq!(listed, vec!["owner/1/a.wasm", "owner/1/b.rego"]);
assert!(backend.list("owner/none").await.unwrap().is_empty());
}
#[tokio::test]
async fn local_backend_rejects_path_traversal_upload() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.put("../escape", b"x").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
let err = backend.put("a/../../b", b"x").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_path_traversal_download() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.get("../etc/passwd").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_path_traversal_exists() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.exists("../etc/passwd").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_path_traversal_delete() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.delete("../sensitive").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_path_traversal_list() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.list("../").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_absolute_path() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.put("/etc/passwd", b"x").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("absolute") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_dot_component() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.put(".", b"x").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("path traversal") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_rejects_empty_key() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
let err = backend.put("", b"x").await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("empty") || msg.contains("resolve key"), "{msg}");
}
#[tokio::test]
async fn local_backend_allows_normal_keys() {
let tmp = tempfile::tempdir().unwrap();
let backend = LocalBackend::new(tmp.path());
backend.put("owners/abc123/cid456", b"x").await.unwrap();
backend.put("objects/bafy123", b"y").await.unwrap();
assert!(backend.exists("owners/abc123/cid456").await.unwrap());
assert!(backend.exists("objects/bafy123").await.unwrap());
}
}