use std::{
cell::RefCell,
collections::{BTreeMap, VecDeque},
rc::Rc,
};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ScenarioBoundary {
BeforeOperation,
AfterDurableCommit,
BeforeResponse,
DuringResponse,
Timeout,
Cancellation,
DroppedConnection,
Cleanup,
ResourceAcquire,
}
impl ScenarioBoundary {
const fn label(self) -> &'static str {
match self {
Self::BeforeOperation => "before-operation",
Self::AfterDurableCommit => "after-durable-commit",
Self::BeforeResponse => "before-response",
Self::DuringResponse => "during-response",
Self::Timeout => "timeout",
Self::Cancellation => "cancellation",
Self::DroppedConnection => "dropped-connection",
Self::Cleanup => "cleanup",
Self::ResourceAcquire => "resource-acquire",
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SimulatorFault {
Timeout,
Cancellation,
DroppedConnection,
CleanupFailure,
ResourceUnavailable,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct FaultPointError;
impl std::fmt::Display for FaultPointError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("fault point must be a bounded ASCII scenario label")
}
}
impl std::error::Error for FaultPointError {}
#[derive(Clone, Debug, Default)]
pub struct FaultInjector {
queues: Rc<RefCell<BTreeMap<String, VecDeque<SimulatorFault>>>>,
}
impl FaultInjector {
pub fn new() -> Self {
Self::default()
}
pub fn inject(&self, point: &str, fault: SimulatorFault) -> Result<(), FaultPointError> {
validate_point(point)?;
self.queues
.borrow_mut()
.entry(point.to_owned())
.or_default()
.push_back(fault);
Ok(())
}
pub fn inject_at(
&self,
operation: &str,
boundary: ScenarioBoundary,
fault: SimulatorFault,
) -> Result<(), FaultPointError> {
self.inject(&format!("{operation}.{}", boundary.label()), fault)
}
pub fn check(&self, point: &str) -> Result<Result<(), SimulatorFault>, FaultPointError> {
validate_point(point)?;
let mut queues = self.queues.borrow_mut();
let fault = queues.get_mut(point).and_then(VecDeque::pop_front);
if queues.get(point).is_some_and(VecDeque::is_empty) {
queues.remove(point);
}
Ok(fault.map_or(Ok(()), Err))
}
pub fn check_at(
&self,
operation: &str,
boundary: ScenarioBoundary,
) -> Result<Result<(), SimulatorFault>, FaultPointError> {
self.check(&format!("{operation}.{}", boundary.label()))
}
pub fn has_pending(&self, point: &str) -> Result<bool, FaultPointError> {
validate_point(point)?;
Ok(self
.queues
.borrow()
.get(point)
.is_some_and(|queue| !queue.is_empty()))
}
}
pub(crate) fn validate_point(point: &str) -> Result<(), FaultPointError> {
if point.is_empty()
|| point.len() > 128
|| !point
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-' | b':'))
{
return Err(FaultPointError);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn faults_are_consumed_once_in_test_selected_order() {
let faults = FaultInjector::new();
faults
.inject(
"oauth.after-durable-commit",
SimulatorFault::DroppedConnection,
)
.unwrap();
faults
.inject("oauth.after-durable-commit", SimulatorFault::Timeout)
.unwrap();
assert!(faults.has_pending("oauth.after-durable-commit").unwrap());
assert_eq!(
faults.check("oauth.after-durable-commit").unwrap(),
Err(SimulatorFault::DroppedConnection)
);
assert_eq!(
faults.check("oauth.after-durable-commit").unwrap(),
Err(SimulatorFault::Timeout)
);
assert_eq!(faults.check("oauth.after-durable-commit").unwrap(), Ok(()));
assert!(!faults.has_pending("oauth.after-durable-commit").unwrap());
}
#[test]
fn canonical_boundaries_do_not_hide_the_operation_identity() {
let faults = FaultInjector::new();
faults
.inject_at(
"oauth.consume",
ScenarioBoundary::AfterDurableCommit,
SimulatorFault::DroppedConnection,
)
.unwrap();
assert_eq!(
faults
.check_at("oauth.consume", ScenarioBoundary::AfterDurableCommit)
.unwrap(),
Err(SimulatorFault::DroppedConnection)
);
assert_eq!(
faults
.check_at("oauth.create", ScenarioBoundary::AfterDurableCommit)
.unwrap(),
Ok(())
);
}
#[test]
fn fault_points_reject_unbounded_or_sensitive_shape() {
let faults = FaultInjector::new();
assert_eq!(
faults.inject("before response", SimulatorFault::Timeout),
Err(FaultPointError)
);
assert_eq!(faults.check(""), Err(FaultPointError));
}
}