use std::{
io::Error as IoError,
ops::Range,
sync::{
Arc, Mutex, MutexGuard,
atomic::{AtomicUsize, Ordering},
},
};
use async_trait::async_trait;
use bytes::Bytes;
use crate::{
runtime_metrics::io::UsageMeter,
storage::{ObjectMeta, StorageError, StorageProvider},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FaultOp {
Head,
Get,
GetRange,
PutAtomic,
PutIfMatch,
Delete,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FaultKind {
Transient,
Precondition,
}
#[derive(Debug)]
struct FaultRule {
op: FaultOp,
kind: FaultKind,
uri_fragment: String,
remaining: usize,
}
#[derive(Debug)]
pub struct FaultStorage {
inner: Arc<dyn StorageProvider>,
rules: Mutex<Vec<FaultRule>>,
fired: AtomicUsize,
}
impl FaultStorage {
pub fn wrap(inner: Arc<dyn StorageProvider>) -> Arc<Self> {
Arc::new(Self {
inner,
rules: Mutex::new(Vec::new()),
fired: AtomicUsize::new(0),
})
}
pub fn fail(&self, op: FaultOp, uri_fragment: &str, times: usize) {
self.fail_with(FaultKind::Transient, op, uri_fragment, times);
}
pub fn fail_with(&self, kind: FaultKind, op: FaultOp, uri_fragment: &str, times: usize) {
self.rules_guard().push(FaultRule {
op,
kind,
uri_fragment: uri_fragment.to_string(),
remaining: times,
});
}
pub fn clear(&self) {
self.rules_guard().clear();
}
pub fn fired(&self) -> usize {
self.fired.load(Ordering::SeqCst)
}
fn rules_guard(&self) -> MutexGuard<'_, Vec<FaultRule>> {
match self.rules.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
}
fn check(&self, op: FaultOp, uri: &str) -> Result<(), StorageError> {
let mut rules = self.rules_guard();
for rule in rules.iter_mut() {
if rule.op == op && rule.remaining > 0 && uri.contains(&rule.uri_fragment) {
rule.remaining -= 1;
self.fired.fetch_add(1, Ordering::SeqCst);
return Err(match rule.kind {
FaultKind::Transient => StorageError::TransientExhausted {
uri: uri.to_string(),
source: Box::new(IoError::other("injected fault")),
},
FaultKind::Precondition => StorageError::PreconditionFailed {
uri: uri.to_string(),
},
});
}
}
Ok(())
}
}
#[async_trait]
impl StorageProvider for FaultStorage {
async fn head(&self, uri: &str) -> Result<ObjectMeta, StorageError> {
self.check(FaultOp::Head, uri)?;
self.inner.head(uri).await
}
async fn get(&self, uri: &str) -> Result<(Bytes, ObjectMeta), StorageError> {
self.check(FaultOp::Get, uri)?;
self.inner.get(uri).await
}
async fn get_range(&self, uri: &str, range: Range<u64>) -> Result<Bytes, StorageError> {
self.check(FaultOp::GetRange, uri)?;
self.inner.get_range(uri, range).await
}
async fn put_atomic(&self, uri: &str, bytes: Bytes) -> Result<Option<String>, StorageError> {
self.check(FaultOp::PutAtomic, uri)?;
self.inner.put_atomic(uri, bytes).await
}
async fn put_if_match(
&self,
uri: &str,
bytes: Bytes,
expected_etag: Option<&str>,
) -> Result<Option<String>, StorageError> {
self.check(FaultOp::PutIfMatch, uri)?;
self.inner.put_if_match(uri, bytes, expected_etag).await
}
async fn put_multipart(
&self,
uri: &str,
) -> Result<Box<dyn object_store::MultipartUpload>, StorageError> {
self.inner.put_multipart(uri).await
}
async fn delete(&self, uri: &str) -> Result<(), StorageError> {
self.check(FaultOp::Delete, uri)?;
self.inner.delete(uri).await
}
async fn list_with_prefix_metadata(
&self,
prefix: &str,
) -> Result<Vec<(String, ObjectMeta)>, StorageError> {
self.inner.list_with_prefix_metadata(prefix).await
}
fn object_store_handle(
&self,
_uri: &str,
) -> Option<(Arc<dyn object_store::ObjectStore>, object_store::path::Path)> {
None
}
fn usage_meter(&self) -> Arc<UsageMeter> {
self.inner.usage_meter()
}
}
#[cfg(test)]
mod tests {
use tempfile::TempDir;
use super::*;
use crate::storage::LocalFsStorageProvider;
#[tokio::test]
async fn rules_burn_down_and_passthroughs_stay_transparent() {
let dir = TempDir::new().expect("tempdir");
let local: Arc<dyn StorageProvider> =
Arc::new(LocalFsStorageProvider::new(dir.path()).expect("local"));
let faults = FaultStorage::wrap(local);
faults
.put_atomic("data/a.bin", Bytes::from_static(b"abc"))
.await
.expect("no rules armed");
faults.fail(FaultOp::Get, "data/", 1);
assert!(faults.get("data/a.bin").await.is_err(), "armed rule fires");
assert_eq!(faults.fired(), 1);
let (bytes, _) = faults.get("data/a.bin").await.expect("rule burned down");
assert_eq!(bytes.as_ref(), b"abc");
faults.fail_with(FaultKind::Precondition, FaultOp::PutIfMatch, "data/", 1);
let err = faults
.put_if_match("data/a.bin", Bytes::from_static(b"xyz"), None)
.await
.expect_err("armed precondition rule fires");
assert!(
matches!(err, StorageError::PreconditionFailed { .. }),
"expected PreconditionFailed, got {err:?}"
);
faults.fail(FaultOp::Get, "data/", 1);
faults.clear();
faults
.get("data/a.bin")
.await
.expect("cleared rules never fire");
assert_eq!(
faults.fired(),
2,
"the get and the precondition rule fired, nothing since"
);
let mut upload = faults.put_multipart("data/m.bin").await.expect("multipart");
upload.abort().await.expect("abort");
assert!(faults.object_store_handle("data/a.bin").is_none());
let _ = faults.usage_meter();
let listed = faults
.list_with_prefix_metadata("data/")
.await
.expect("listing delegates");
assert_eq!(listed.len(), 1, "only the seeded object");
}
}