macro_rules! fault_point {
($stack:expr, $op:expr) => {
#[cfg(all(debug_assertions, feature = "fault-injection"))]
{
if let Some(__fault_err) = $crate::fault::__consult(&$stack.fault, $op) {
return Err(__fault_err);
}
}
};
}
pub(crate) use fault_point;
#[cfg(all(debug_assertions, feature = "fault-injection"))]
pub use active::*;
#[cfg(all(debug_assertions, feature = "fault-injection"))]
mod active {
use std::io;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
pub trait FaultPolicy: Send + Sync {
fn next_fault(&self, op: &'static str, seq: u64) -> Option<io::Error>;
}
#[derive(Default)]
pub struct FaultState {
policy: Mutex<Option<Arc<dyn FaultPolicy>>>,
seq: AtomicU64,
}
impl FaultState {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn set(&self, policy: Option<Arc<dyn FaultPolicy>>) {
let mut guard = self.policy.lock().unwrap();
*guard = policy;
self.seq.store(0, Ordering::SeqCst);
}
pub(crate) fn get(&self) -> Option<Arc<dyn FaultPolicy>> {
self.policy.lock().unwrap().clone()
}
}
impl std::fmt::Debug for FaultState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FaultState")
.field("armed", &self.policy.lock().unwrap().is_some())
.field("seq", &self.seq.load(Ordering::Relaxed))
.finish()
}
}
#[inline]
pub(crate) fn __consult(state: &FaultState, op: &'static str) -> Option<io::Error> {
let policy = state.get()?;
let seq = state.seq.fetch_add(1, Ordering::SeqCst);
policy.next_fault(op, seq)
}
}
#[cfg(all(test, debug_assertions, feature = "fault-injection"))]
mod tests {
use super::FaultPolicy;
use crate::BStack;
use std::io;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
fn mk_stack() -> (BStack, std::path::PathBuf) {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let pid = std::process::id();
let path = std::env::temp_dir().join(format!("bstack_fault_{pid}_{id}.bin"));
let stack = BStack::open(&path).unwrap();
(stack, path)
}
struct Guard(std::path::PathBuf);
impl Drop for Guard {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.0);
}
}
struct FailOpAt {
op: &'static str,
at: u64,
kind: io::ErrorKind,
}
impl FaultPolicy for FailOpAt {
fn next_fault(&self, op: &'static str, seq: u64) -> Option<io::Error> {
(op == self.op && seq == self.at)
.then(|| io::Error::new(self.kind, format!("injected at {op}#{seq}")))
}
}
struct FailAll {
seen: std::sync::Mutex<Vec<(&'static str, u64)>>,
}
impl FaultPolicy for FailAll {
fn next_fault(&self, op: &'static str, seq: u64) -> Option<io::Error> {
self.seen.lock().unwrap().push((op, seq));
Some(io::Error::new(io::ErrorKind::Other, "always"))
}
}
#[test]
fn injects_then_rolls_back_and_disarms() {
let (stack, path) = mk_stack();
let _g = Guard(path);
stack.push(b"seed").unwrap();
assert_eq!(stack.len().unwrap(), 4);
stack.set_fault_policy(Some(Arc::new(FailOpAt {
op: "push",
at: 0,
kind: io::ErrorKind::StorageFull,
})));
let err = stack.push(b"more").unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::StorageFull);
assert_eq!(stack.len().unwrap(), 4);
stack.set_fault_policy(None);
stack.push(b"more").unwrap();
assert_eq!(stack.len().unwrap(), 8);
drop(stack);
}
#[test]
fn validation_beats_injected_fault() {
let (stack, path) = mk_stack();
let _g = Guard(path);
stack.push(b"abc").unwrap();
let policy = Arc::new(FailAll {
seen: std::sync::Mutex::new(Vec::new()),
});
stack.set_fault_policy(Some(policy.clone()));
let err = stack.pop(999).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
assert!(
!policy
.seen
.lock()
.unwrap()
.iter()
.any(|(op, _)| *op == "pop")
);
let err = stack.pop(1).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Other);
assert!(
policy
.seen
.lock()
.unwrap()
.iter()
.any(|(op, _)| *op == "pop")
);
drop(stack);
}
#[test]
fn sequence_counter_is_shared_and_monotonic() {
let (stack, path) = mk_stack();
let _g = Guard(path);
let policy = Arc::new(FailAll {
seen: std::sync::Mutex::new(Vec::new()),
});
stack.set_fault_policy(Some(policy.clone()));
let _ = stack.len();
let _ = stack.push(b"x");
let _ = stack.len();
let seen = policy.seen.lock().unwrap();
assert_eq!(&seen[..], &[("len", 0), ("push", 1), ("len", 2)]);
drop(seen);
drop(stack);
}
#[test]
fn re_arming_resets_the_counter() {
let (stack, path) = mk_stack();
let _g = Guard(path);
let policy = Arc::new(FailAll {
seen: std::sync::Mutex::new(Vec::new()),
});
stack.set_fault_policy(Some(policy.clone()));
let _ = stack.len();
let _ = stack.len();
let policy2 = Arc::new(FailAll {
seen: std::sync::Mutex::new(Vec::new()),
});
stack.set_fault_policy(Some(policy2.clone()));
let _ = stack.len();
assert_eq!(&policy2.seen.lock().unwrap()[..], &[("len", 0)]);
drop(stack);
}
}