use anyhow::Result;
use lora_wal::{WalBufferedCommitError, WalRecorder};
#[derive(Debug, Clone, Copy)]
pub(crate) enum WalAbortPolicy {
AbortOnly,
PoisonIfMutated(&'static str),
}
pub(crate) struct WalWriteScope<'a> {
recorder: &'a WalRecorder,
abort_policy: WalAbortPolicy,
}
impl<'a> WalWriteScope<'a> {
pub(crate) fn arm(recorder: &'a WalRecorder, abort_policy: WalAbortPolicy) -> Result<Self> {
recorder.arm().map_err(WalBufferedCommitError::Arm)?;
Ok(Self {
recorder,
abort_policy,
})
}
pub(crate) fn finish<R>(self, result: &Result<R>) -> Result<bool> {
let wrote_commit = match result {
Ok(_) => self.recorder.commit()?.wrote(),
Err(_) => {
abort_armed(self.recorder, self.abort_policy)?;
false
}
};
ensure_wal_not_poisoned(self.recorder)?;
Ok(wrote_commit)
}
}
pub(crate) fn ensure_wal_not_poisoned(recorder: &WalRecorder) -> Result<()> {
if let Some(reason) = recorder.poisoned_reason() {
return Err(WalBufferedCommitError::Poisoned(reason).into());
}
Ok(())
}
pub(crate) fn ensure_wal_query_can_start(recorder: &WalRecorder) -> Result<()> {
if let Some(reason) = recorder.poisoned_reason() {
return Err(WalBufferedCommitError::Poisoned(reason).into());
}
Ok(())
}
fn abort_armed(recorder: &WalRecorder, policy: WalAbortPolicy) -> Result<()> {
let aborted_after_mutation = matches!(recorder.abort(), Ok(true));
if aborted_after_mutation {
if let WalAbortPolicy::PoisonIfMutated(reason) = policy {
recorder.poison(reason);
}
}
Ok(())
}