use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum Backoff {
Fixed,
Linear,
Exponential,
}
impl Backoff {
pub fn delay_seconds(self, base: u64, attempt: u32) -> u64 {
let n = attempt.max(1);
match self {
Backoff::Fixed => base,
Backoff::Linear => base.saturating_mul(n as u64),
Backoff::Exponential => base.saturating_mul(1u64 << (n - 1).min(16)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "action", rename_all = "snake_case", deny_unknown_fields)]
pub enum RecoveryAction {
Retry {
#[serde(default = "default_attempts")]
max_attempts: u32,
#[serde(default = "default_base_delay")]
base_delay_seconds: u64,
#[serde(default = "default_backoff")]
backoff: Backoff,
},
Revise {
#[serde(default = "default_attempts")]
max_attempts: u32,
},
Fallback {},
Escalate {},
Pause {},
RestoreCheckpoint {},
Stop {},
}
fn default_attempts() -> u32 {
3
}
fn default_base_delay() -> u64 {
2
}
fn default_backoff() -> Backoff {
Backoff::Exponential
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct Recovery {
#[serde(default = "default_transient")]
pub transient_error: RecoveryAction,
#[serde(default = "default_invalid_output")]
pub invalid_output: RecoveryAction,
#[serde(default = "default_tool_unavailable")]
pub tool_unavailable: RecoveryAction,
#[serde(default = "default_repeated_failure")]
pub repeated_failure: RecoveryAction,
#[serde(default = "default_safety_violation")]
pub safety_violation: RecoveryAction,
#[serde(default = "default_resource_exhaustion")]
pub resource_exhaustion: RecoveryAction,
#[serde(default = "default_corrupted_state")]
pub corrupted_state: RecoveryAction,
}
fn default_transient() -> RecoveryAction {
RecoveryAction::Retry {
max_attempts: default_attempts(),
base_delay_seconds: default_base_delay(),
backoff: Backoff::Exponential,
}
}
fn default_invalid_output() -> RecoveryAction {
RecoveryAction::Revise { max_attempts: 2 }
}
fn default_tool_unavailable() -> RecoveryAction {
RecoveryAction::Fallback {}
}
fn default_repeated_failure() -> RecoveryAction {
RecoveryAction::Escalate {}
}
fn default_safety_violation() -> RecoveryAction {
RecoveryAction::Stop {}
}
fn default_resource_exhaustion() -> RecoveryAction {
RecoveryAction::Pause {}
}
fn default_corrupted_state() -> RecoveryAction {
RecoveryAction::RestoreCheckpoint {}
}
impl Default for Recovery {
fn default() -> Self {
Self {
transient_error: default_transient(),
invalid_output: default_invalid_output(),
tool_unavailable: default_tool_unavailable(),
repeated_failure: default_repeated_failure(),
safety_violation: default_safety_violation(),
resource_exhaustion: default_resource_exhaustion(),
corrupted_state: default_corrupted_state(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum FailureClass {
TransientError,
InvalidOutput,
ToolUnavailable,
RepeatedFailure,
SafetyViolation,
ResourceExhaustion,
CorruptedState,
}
impl FailureClass {
pub const ALL: [FailureClass; 7] = [
FailureClass::TransientError,
FailureClass::InvalidOutput,
FailureClass::ToolUnavailable,
FailureClass::RepeatedFailure,
FailureClass::SafetyViolation,
FailureClass::ResourceExhaustion,
FailureClass::CorruptedState,
];
pub fn key(self) -> &'static str {
match self {
FailureClass::TransientError => "transient_error",
FailureClass::InvalidOutput => "invalid_output",
FailureClass::ToolUnavailable => "tool_unavailable",
FailureClass::RepeatedFailure => "repeated_failure",
FailureClass::SafetyViolation => "safety_violation",
FailureClass::ResourceExhaustion => "resource_exhaustion",
FailureClass::CorruptedState => "corrupted_state",
}
}
}
impl Recovery {
pub fn action_for(&self, class: FailureClass) -> RecoveryAction {
match class {
FailureClass::TransientError => self.transient_error,
FailureClass::InvalidOutput => self.invalid_output,
FailureClass::ToolUnavailable => self.tool_unavailable,
FailureClass::RepeatedFailure => self.repeated_failure,
FailureClass::SafetyViolation => self.safety_violation,
FailureClass::ResourceExhaustion => self.resource_exhaustion,
FailureClass::CorruptedState => self.corrupted_state,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn exponential_backoff_cannot_overflow_into_a_short_delay() {
let big = Backoff::Exponential.delay_seconds(2, 4096);
let small = Backoff::Exponential.delay_seconds(2, 4);
assert!(big >= small, "{big} should not be shorter than {small}");
}
#[test]
fn the_first_retry_waits_the_base_delay_under_every_strategy() {
for b in [Backoff::Fixed, Backoff::Linear, Backoff::Exponential] {
assert_eq!(b.delay_seconds(3, 1), 3, "{b:?}");
}
assert_eq!(Backoff::Exponential.delay_seconds(3, 2), 6);
assert_eq!(Backoff::Exponential.delay_seconds(3, 3), 12);
}
#[test]
fn every_class_is_named_as_the_config_spells_it() {
for class in FailureClass::ALL {
let json = serde_json::to_string(&class).unwrap();
assert_eq!(json.trim_matches('"'), class.key());
}
}
#[test]
fn a_safety_violation_stops_rather_than_retries() {
assert_eq!(
Recovery::default().action_for(FailureClass::SafetyViolation),
RecoveryAction::Stop {}
);
}
}