use std::collections::BTreeMap;
use std::error::Error;
use std::fmt;
use std::sync::{Arc, Mutex, PoisonError};
use serde_json::{Map, Number, Value};
use sha2::{Digest, Sha256};
use crate::{
BackoffSleeper, BoxFuture, ChunkDeliveryMode, ChunkTransactionContext, ClassifierRevision,
FailureCategory, FaultPhase, FaultPolicy, RetryLimit, RetryOrdinal, RetryStateLimit,
SkipCounts, StepName,
};
const RETRY_KEY_DOMAIN: &[u8] = b"oxide-batch/retry-key/1";
#[derive(Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct RetryKey([u8; 32]);
impl RetryKey {
#[must_use]
pub(crate) fn derive(
definition_digest: &[u8; 32],
step_name: &StepName,
phase: FaultPhase,
checkpoint_digest: &[u8; 32],
ordinal: u64,
) -> Self {
let mut hasher = Sha256::new();
hasher.update(RETRY_KEY_DOMAIN);
hasher.update(definition_digest);
hasher.update((step_name.as_str().len() as u64).to_be_bytes());
hasher.update(step_name.as_str().as_bytes());
hasher.update(phase.as_str().as_bytes());
hasher.update([0]);
hasher.update(checkpoint_digest);
hasher.update(ordinal.to_be_bytes());
Self(hasher.finalize().into())
}
#[must_use]
pub const fn from_bytes(digest: [u8; 32]) -> Self {
Self(digest)
}
#[must_use]
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
impl fmt::Debug for RetryKey {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RetryKey")
.field("digest", &"<redacted>")
.finish()
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct RetryReservation {
key: RetryKey,
phase: FaultPhase,
category: FailureCategory,
ordinal: RetryOrdinal,
}
impl RetryReservation {
#[must_use]
pub const fn new(
key: RetryKey,
phase: FaultPhase,
category: FailureCategory,
ordinal: RetryOrdinal,
) -> Self {
Self {
key,
phase,
category,
ordinal,
}
}
#[must_use]
pub const fn key(self) -> RetryKey {
self.key
}
#[must_use]
pub const fn phase(self) -> FaultPhase {
self.phase
}
#[must_use]
pub const fn category(self) -> FailureCategory {
self.category
}
#[must_use]
pub const fn ordinal(self) -> RetryOrdinal {
self.ordinal
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum FaultStateError {
CapacityExhausted {
max: u32,
},
StaleReservation,
Corrupt(FaultStateFormatError),
Unbound,
Unavailable,
}
impl fmt::Display for FaultStateError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::CapacityExhausted { max } => {
write!(
formatter,
"step already retains {max} unresolved retry keys"
)
}
Self::StaleReservation => {
formatter.write_str("retry reservation lost to a newer persisted ordinal")
}
Self::Corrupt(error) => write!(formatter, "durable fault state is unusable: {error}"),
Self::Unbound => formatter.write_str("durable fault state has no bound step execution"),
Self::Unavailable => formatter.write_str("fault state is unavailable"),
}
}
}
impl Error for FaultStateError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
Self::Corrupt(error) => Some(error),
_ => None,
}
}
}
impl From<FaultStateFormatError> for FaultStateError {
fn from(error: FaultStateFormatError) -> Self {
Self::Corrupt(error)
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct FaultStateEntry {
key: RetryKey,
phase: FaultPhase,
category: FailureCategory,
ordinal: RetryOrdinal,
revision: ClassifierRevision,
}
impl FaultStateEntry {
#[must_use]
pub const fn new(
key: RetryKey,
phase: FaultPhase,
category: FailureCategory,
ordinal: RetryOrdinal,
revision: ClassifierRevision,
) -> Self {
Self {
key,
phase,
category,
ordinal,
revision,
}
}
#[must_use]
pub const fn key(&self) -> RetryKey {
self.key
}
#[must_use]
pub const fn phase(&self) -> FaultPhase {
self.phase
}
#[must_use]
pub const fn category(&self) -> FailureCategory {
self.category
}
#[must_use]
pub const fn ordinal(&self) -> RetryOrdinal {
self.ordinal
}
#[must_use]
pub const fn revision(&self) -> &ClassifierRevision {
&self.revision
}
fn to_json(&self) -> Value {
let mut object = Map::new();
object.insert(
String::from("category"),
Value::String(String::from(self.category.durable_code())),
);
object.insert(String::from("key"), Value::String(hex(self.key.as_bytes())));
object.insert(
String::from("ordinal"),
Value::Number(Number::from(self.ordinal.get())),
);
object.insert(
String::from("phase"),
Value::String(String::from(self.phase.as_str())),
);
object.insert(
String::from("revision"),
Value::String(String::from(self.revision.as_str())),
);
Value::Object(object)
}
fn from_json(value: &Value) -> Result<Self, FaultStateFormatError> {
let object = value
.as_object()
.ok_or(FaultStateFormatError::MalformedEntry)?;
let key = object
.get("key")
.and_then(Value::as_str)
.and_then(unhex)
.map(RetryKey::from_bytes)
.ok_or(FaultStateFormatError::MalformedEntry)?;
let phase = object
.get("phase")
.and_then(Value::as_str)
.and_then(FaultPhase::from_durable_name)
.ok_or(FaultStateFormatError::UnknownEnumeration)?;
let category = object
.get("category")
.and_then(Value::as_str)
.and_then(FailureCategory::from_durable_code)
.ok_or(FaultStateFormatError::UnknownEnumeration)?;
let ordinal = object
.get("ordinal")
.and_then(Value::as_u64)
.and_then(|value| u32::try_from(value).ok())
.and_then(|value| RetryOrdinal::new(value).ok())
.ok_or(FaultStateFormatError::MalformedEntry)?;
let revision = object
.get("revision")
.and_then(Value::as_str)
.and_then(|value| ClassifierRevision::new(value).ok())
.ok_or(FaultStateFormatError::MalformedEntry)?;
Ok(Self::new(key, phase, category, ordinal, revision))
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct FaultStateEnvelope {
checkpoint_digest: [u8; 32],
entries: Vec<FaultStateEntry>,
}
impl FaultStateEnvelope {
pub const FORMAT: &'static str = "oxide-batch.fault-state";
pub const FORMAT_VERSION: u16 = 1;
pub const SCHEMA_VERSION: u32 = 1;
pub const MAX_BYTES: usize = 64 * 1024;
pub const MAX_ENTRIES: usize = 256;
#[must_use]
pub const fn empty() -> Self {
Self {
checkpoint_digest: [0; 32],
entries: Vec::new(),
}
}
pub fn new(
checkpoint_digest: [u8; 32],
entries: impl IntoIterator<Item = FaultStateEntry>,
) -> Result<Self, FaultStateFormatError> {
let mut entries: Vec<FaultStateEntry> = entries.into_iter().collect();
if entries.len() > Self::MAX_ENTRIES {
return Err(FaultStateFormatError::TooManyEntries {
max: Self::MAX_ENTRIES,
});
}
entries.sort_by_key(FaultStateEntry::key);
if entries.windows(2).any(|pair| pair[0].key == pair[1].key) {
return Err(FaultStateFormatError::DuplicateKey);
}
if entries.is_empty() {
if checkpoint_digest != [0; 32] {
return Err(FaultStateFormatError::CheckpointMismatch);
}
} else if checkpoint_digest == [0; 32] {
return Err(FaultStateFormatError::CheckpointMismatch);
}
Ok(Self {
checkpoint_digest,
entries,
})
}
#[must_use]
pub const fn checkpoint_digest(&self) -> &[u8; 32] {
&self.checkpoint_digest
}
#[must_use]
pub fn entries(&self) -> &[FaultStateEntry] {
&self.entries
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn reserved_ordinal(&self, key: RetryKey) -> Option<RetryOrdinal> {
self.entry(key).map(FaultStateEntry::ordinal)
}
#[must_use]
pub fn entry(&self, key: RetryKey) -> Option<&FaultStateEntry> {
self.entries
.binary_search_by(|entry| entry.key.cmp(&key))
.ok()
.map(|index| &self.entries[index])
}
pub fn reserved(
&self,
entry: FaultStateEntry,
checkpoint_digest: [u8; 32],
limit: RetryStateLimit,
) -> Result<Self, FaultStateError> {
if !self.entries.is_empty() && self.checkpoint_digest != checkpoint_digest {
return Err(FaultStateError::Corrupt(
FaultStateFormatError::CheckpointMismatch,
));
}
let expected = self
.reserved_ordinal(entry.key())
.unwrap_or(RetryOrdinal::INITIAL)
.checked_next()
.map_err(|_| FaultStateError::StaleReservation)?;
if entry.ordinal() != expected {
return Err(FaultStateError::StaleReservation);
}
let mut entries = self.entries.clone();
match entries.binary_search_by(|existing| existing.key.cmp(&entry.key())) {
Ok(index) => entries[index] = entry,
Err(index) => {
if entries.len() >= limit.get() as usize {
return Err(FaultStateError::CapacityExhausted { max: limit.get() });
}
entries.insert(index, entry);
}
}
Ok(Self {
checkpoint_digest,
entries,
})
}
pub fn to_canonical_json(&self) -> Result<Vec<u8>, FaultStateFormatError> {
let mut object = Map::new();
object.insert(
String::from("checkpoint"),
Value::String(hex(&self.checkpoint_digest)),
);
object.insert(
String::from("entries"),
Value::Array(self.entries.iter().map(FaultStateEntry::to_json).collect()),
);
let bytes = serde_json::to_vec(&Value::Object(object))
.map_err(|_| FaultStateFormatError::Malformed)?;
if bytes.len() > Self::MAX_BYTES {
return Err(FaultStateFormatError::TooLarge {
max_bytes: Self::MAX_BYTES,
});
}
Ok(bytes)
}
pub fn checksum(&self) -> Result<[u8; 32], FaultStateFormatError> {
Ok(Sha256::digest(self.to_canonical_json()?).into())
}
pub fn from_canonical_json(
format_version: u16,
schema: &str,
schema_version: u32,
bytes: &[u8],
checksum: &[u8; 32],
) -> Result<Self, FaultStateFormatError> {
if format_version != Self::FORMAT_VERSION || schema != Self::FORMAT {
return Err(FaultStateFormatError::UnsupportedFormat);
}
if schema_version != Self::SCHEMA_VERSION {
return Err(FaultStateFormatError::UnsupportedSchemaVersion);
}
if bytes.len() > Self::MAX_BYTES {
return Err(FaultStateFormatError::TooLarge {
max_bytes: Self::MAX_BYTES,
});
}
let value: Value =
serde_json::from_slice(bytes).map_err(|_| FaultStateFormatError::Malformed)?;
let object = value.as_object().ok_or(FaultStateFormatError::Malformed)?;
let checkpoint_digest = object
.get("checkpoint")
.and_then(Value::as_str)
.and_then(unhex)
.ok_or(FaultStateFormatError::Malformed)?;
let raw = object
.get("entries")
.and_then(Value::as_array)
.ok_or(FaultStateFormatError::Malformed)?;
if raw.len() > Self::MAX_ENTRIES {
return Err(FaultStateFormatError::TooManyEntries {
max: Self::MAX_ENTRIES,
});
}
let entries = raw
.iter()
.map(FaultStateEntry::from_json)
.collect::<Result<Vec<_>, _>>()?;
if entries.windows(2).any(|pair| pair[0].key >= pair[1].key) {
return Err(FaultStateFormatError::UnsortedEntries);
}
let envelope = Self::new(checkpoint_digest, entries)?;
if &envelope.checksum()? != checksum {
return Err(FaultStateFormatError::ChecksumMismatch);
}
Ok(envelope)
}
pub fn validate_for(
&self,
retry_limit: RetryLimit,
state_limit: RetryStateLimit,
checkpoint_digest: &[u8; 32],
) -> Result<(), FaultStateFormatError> {
if self.entries.len() > state_limit.get() as usize {
return Err(FaultStateFormatError::TooManyEntries {
max: state_limit.get() as usize,
});
}
if self
.entries
.iter()
.any(|entry| entry.ordinal().get() > retry_limit.get())
{
return Err(FaultStateFormatError::OrdinalAboveLimit {
max: retry_limit.get(),
});
}
if !self.entries.is_empty() && &self.checkpoint_digest != checkpoint_digest {
return Err(FaultStateFormatError::CheckpointMismatch);
}
Ok(())
}
}
impl Default for FaultStateEnvelope {
fn default() -> Self {
Self::empty()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum FaultStateFormatError {
UnsupportedFormat,
UnsupportedSchemaVersion,
Malformed,
MalformedEntry,
UnknownEnumeration,
ChecksumMismatch,
TooLarge {
max_bytes: usize,
},
TooManyEntries {
max: usize,
},
DuplicateKey,
UnsortedEntries,
OrdinalAboveLimit {
max: u32,
},
CheckpointMismatch,
}
impl fmt::Display for FaultStateFormatError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::UnsupportedFormat => formatter.write_str("fault-state format is unsupported"),
Self::UnsupportedSchemaVersion => {
formatter.write_str("fault-state schema version is unsupported")
}
Self::Malformed => formatter.write_str("fault state is malformed"),
Self::MalformedEntry => formatter.write_str("fault-state entry is malformed"),
Self::UnknownEnumeration => {
formatter.write_str("fault state contains an unknown enumeration value")
}
Self::ChecksumMismatch => formatter.write_str("fault-state checksum does not match"),
Self::TooLarge { max_bytes } => {
write!(formatter, "fault state exceeds {max_bytes} bytes")
}
Self::TooManyEntries { max } => {
write!(formatter, "fault state retains more than {max} keys")
}
Self::DuplicateKey => formatter.write_str("fault state repeats one retry key"),
Self::UnsortedEntries => formatter.write_str("fault state is not digest-sorted"),
Self::OrdinalAboveLimit { max } => {
write!(formatter, "fault state retains an ordinal above {max}")
}
Self::CheckpointMismatch => {
formatter.write_str("fault state belongs to a superseded checkpoint")
}
}
}
}
impl Error for FaultStateFormatError {}
fn hex(bytes: &[u8; 32]) -> String {
let mut text = String::with_capacity(64);
for byte in bytes {
text.push(char::from_digit(u32::from(byte >> 4), 16).unwrap_or('0'));
text.push(char::from_digit(u32::from(byte & 0x0f), 16).unwrap_or('0'));
}
text
}
fn unhex(text: &str) -> Option<[u8; 32]> {
if text.len() != 64 {
return None;
}
let mut bytes = [0_u8; 32];
let raw = text.as_bytes();
for (index, slot) in bytes.iter_mut().enumerate() {
let high = char::from(raw[index * 2]).to_digit(16)?;
let low = char::from(raw[index * 2 + 1]).to_digit(16)?;
*slot = u8::try_from(high * 16 + low).ok()?;
}
Some(bytes)
}
pub trait FaultStateStore: Send + Sync {
fn bind(
&self,
_context: ChunkTransactionContext,
) -> BoxFuture<'_, Result<(), FaultStateError>> {
Box::pin(std::future::ready(Ok(())))
}
fn reserved_ordinal(
&self,
key: RetryKey,
) -> BoxFuture<'_, Result<Option<RetryOrdinal>, FaultStateError>>;
fn reserve(&self, reservation: RetryReservation) -> BoxFuture<'_, Result<(), FaultStateError>>;
fn resolve(&self, key: RetryKey) -> BoxFuture<'_, Result<(), FaultStateError>>;
fn clear_resolved(&self) -> BoxFuture<'_, Result<(), FaultStateError>>;
fn unresolved(&self) -> BoxFuture<'_, Result<u32, FaultStateError>>;
}
#[derive(Clone, Copy, Debug)]
struct RetryEntry {
ordinal: RetryOrdinal,
resolved: bool,
}
#[derive(Debug)]
pub struct InMemoryFaultState {
limit: RetryStateLimit,
entries: Mutex<BTreeMap<RetryKey, RetryEntry>>,
}
impl InMemoryFaultState {
#[must_use]
pub fn new(limit: RetryStateLimit) -> Self {
Self {
limit,
entries: Mutex::new(BTreeMap::new()),
}
}
fn with_entries<T>(&self, body: impl FnOnce(&mut BTreeMap<RetryKey, RetryEntry>) -> T) -> T {
let mut entries = self.entries.lock().unwrap_or_else(PoisonError::into_inner);
body(&mut entries)
}
}
impl FaultStateStore for InMemoryFaultState {
fn reserved_ordinal(
&self,
key: RetryKey,
) -> BoxFuture<'_, Result<Option<RetryOrdinal>, FaultStateError>> {
let result = self.with_entries(|entries| entries.get(&key).map(|entry| entry.ordinal));
Box::pin(std::future::ready(Ok(result)))
}
fn reserve(&self, reservation: RetryReservation) -> BoxFuture<'_, Result<(), FaultStateError>> {
let limit = self.limit;
let result = self.with_entries(|entries| {
let expected = entries
.get(&reservation.key())
.map_or(RetryOrdinal::INITIAL, |entry| entry.ordinal)
.checked_next()
.map_err(|_| FaultStateError::StaleReservation)?;
if reservation.ordinal() != expected {
return Err(FaultStateError::StaleReservation);
}
let unresolved = entries.values().filter(|entry| !entry.resolved).count();
let is_new = !entries.contains_key(&reservation.key());
if is_new && unresolved >= limit.get() as usize {
return Err(FaultStateError::CapacityExhausted { max: limit.get() });
}
entries.insert(
reservation.key(),
RetryEntry {
ordinal: reservation.ordinal(),
resolved: false,
},
);
Ok(())
});
Box::pin(std::future::ready(result))
}
fn resolve(&self, key: RetryKey) -> BoxFuture<'_, Result<(), FaultStateError>> {
self.with_entries(|entries| {
if let Some(entry) = entries.get_mut(&key) {
entry.resolved = true;
}
});
Box::pin(std::future::ready(Ok(())))
}
fn clear_resolved(&self) -> BoxFuture<'_, Result<(), FaultStateError>> {
self.with_entries(|entries| entries.retain(|_, entry| !entry.resolved));
Box::pin(std::future::ready(Ok(())))
}
fn unresolved(&self) -> BoxFuture<'_, Result<u32, FaultStateError>> {
let count = self.with_entries(|entries| entries.values().filter(|e| !e.resolved).count());
let result = u32::try_from(count).map_err(|_| FaultStateError::Unavailable);
Box::pin(std::future::ready(result))
}
}
#[derive(Clone)]
pub struct FaultRuntime {
policy: Arc<FaultPolicy>,
sleeper: Arc<dyn BackoffSleeper>,
state: Arc<dyn FaultStateStore>,
delivery_mode: ChunkDeliveryMode,
}
impl FaultRuntime {
pub fn new(
policy: FaultPolicy,
sleeper: Arc<dyn BackoffSleeper>,
state: Arc<dyn FaultStateStore>,
delivery_mode: ChunkDeliveryMode,
) -> Result<Self, crate::FaultPolicyError> {
policy.validate_capabilities(matches!(
delivery_mode,
ChunkDeliveryMode::AtomicSameResource
))?;
Ok(Self {
policy: Arc::new(policy),
sleeper,
state,
delivery_mode,
})
}
#[must_use]
pub fn policy(&self) -> &FaultPolicy {
&self.policy
}
#[must_use]
pub fn sleeper(&self) -> &dyn BackoffSleeper {
self.sleeper.as_ref()
}
#[must_use]
pub fn state(&self) -> &dyn FaultStateStore {
self.state.as_ref()
}
#[must_use]
pub const fn delivery_mode(&self) -> ChunkDeliveryMode {
self.delivery_mode
}
}
impl fmt::Debug for FaultRuntime {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("FaultRuntime")
.field("retry_limit", &self.policy.retry_limit())
.field("retry_state_limit", &self.policy.retry_state_limit())
.field("skip_limit", &self.policy.skip_limit())
.field("backoff", &self.policy.backoff().kind())
.field("delivery_mode", &self.delivery_mode)
.finish_non_exhaustive()
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub struct FaultProgress {
retries: RetryCounts,
skips: SkipCounts,
rollbacks: u64,
no_rollbacks: u64,
}
impl FaultProgress {
pub const NONE: Self = Self {
retries: RetryCounts::ZERO,
skips: SkipCounts::ZERO,
rollbacks: 0,
no_rollbacks: 0,
};
#[must_use]
pub const fn new(
retries: RetryCounts,
skips: SkipCounts,
rollbacks: u64,
no_rollbacks: u64,
) -> Self {
Self {
retries,
skips,
rollbacks,
no_rollbacks,
}
}
#[must_use]
pub const fn retries(self) -> RetryCounts {
self.retries
}
#[must_use]
pub const fn skips(self) -> SkipCounts {
self.skips
}
#[must_use]
pub const fn rollbacks(self) -> u64 {
self.rollbacks
}
#[must_use]
pub const fn no_rollbacks(self) -> u64 {
self.no_rollbacks
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct RetryCounts {
read: u64,
process: u64,
write: u64,
}
impl RetryCounts {
pub const ZERO: Self = Self {
read: 0,
process: 0,
write: 0,
};
#[must_use]
pub const fn new(read: u64, process: u64, write: u64) -> Self {
Self {
read,
process,
write,
}
}
#[must_use]
pub const fn read(self) -> u64 {
self.read
}
#[must_use]
pub const fn process(self) -> u64 {
self.process
}
#[must_use]
pub const fn write(self) -> u64 {
self.write
}
#[must_use]
pub const fn increment(mut self, phase: FaultPhase) -> Self {
let counter = match phase {
FaultPhase::Read => &mut self.read,
FaultPhase::Process => &mut self.process,
FaultPhase::Write => &mut self.write,
_ => return self,
};
*counter = counter.saturating_add(1);
self
}
}