use std::str::FromStr;
pub const FAULT_POINT_ENV: &str = "MIDENUP_FAULT_POINT";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FaultPoint {
PostPrepare,
PostStage,
PostVerify,
PostCommit,
PostRecord,
PostDerive,
PostUninstallTombstone,
PreMigrationCommit,
}
impl FaultPoint {
pub const ALL: [FaultPoint; 8] = [
Self::PostPrepare,
Self::PostStage,
Self::PostVerify,
Self::PostCommit,
Self::PostRecord,
Self::PostDerive,
Self::PostUninstallTombstone,
Self::PreMigrationCommit,
];
pub const PUBLICATION: [FaultPoint; 6] = [
Self::PostPrepare,
Self::PostStage,
Self::PostVerify,
Self::PostCommit,
Self::PostRecord,
Self::PostDerive,
];
pub const fn as_str(self) -> &'static str {
match self {
Self::PostPrepare => "post-prepare",
Self::PostStage => "post-stage",
Self::PostVerify => "post-verify",
Self::PostCommit => "post-commit",
Self::PostRecord => "post-record",
Self::PostDerive => "post-derive",
Self::PostUninstallTombstone => "post-uninstall-tombstone",
Self::PreMigrationCommit => "pre-migration-commit",
}
}
}
impl std::fmt::Display for FaultPoint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for FaultPoint {
type Err = String;
fn from_str(value: &str) -> Result<Self, Self::Err> {
FaultPoint::ALL
.into_iter()
.find(|point| point.as_str() == value)
.ok_or_else(|| format!("unknown fault point '{value}'"))
}
}
#[derive(Debug, thiserror::Error)]
#[error("aborted at the injected fault point '{point}'")]
pub struct InjectedFault {
pub point: FaultPoint,
}
#[cfg(feature = "fault-injection")]
pub fn fail_at(point: FaultPoint) -> Result<(), InjectedFault> {
let armed = std::env::var(FAULT_POINT_ENV).ok();
match armed.as_deref().map(FaultPoint::from_str) {
Some(Ok(armed)) if armed == point => Err(InjectedFault { point }),
_ => Ok(()),
}
}
#[cfg(not(feature = "fault-injection"))]
#[inline(always)]
pub fn fail_at(_point: FaultPoint) -> Result<(), InjectedFault> {
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_point_round_trips_through_its_name() {
for point in FaultPoint::ALL {
assert_eq!(point.as_str().parse::<FaultPoint>().unwrap(), point);
}
assert!("post-nothing".parse::<FaultPoint>().is_err());
}
#[cfg(not(feature = "fault-injection"))]
#[test]
fn faults_cannot_be_armed_in_a_release_build() {
for point in FaultPoint::ALL {
assert!(fail_at(point).is_ok());
}
}
}