use crate::error::{SageError, SageResult};
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Strategy {
#[default]
OneForOne,
OneForAll,
RestForOne,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum RestartPolicy {
#[default]
Permanent,
Transient,
Temporary,
}
#[derive(Debug, Clone)]
pub struct RestartConfig {
pub max_restarts: u32,
pub within: Duration,
}
impl Default for RestartConfig {
fn default() -> Self {
Self {
max_restarts: 5,
within: Duration::from_secs(60),
}
}
}
#[cfg(not(target_arch = "wasm32"))]
mod native {
use super::*;
use std::collections::VecDeque;
use std::future::Future;
use std::pin::Pin;
use std::time::Instant;
use tokio::task::JoinHandle;
struct RestartTracker {
timestamps: VecDeque<Instant>,
config: RestartConfig,
}
impl RestartTracker {
fn new(config: RestartConfig) -> Self {
Self {
timestamps: VecDeque::new(),
config,
}
}
fn record_restart(&mut self) -> bool {
let now = Instant::now();
while let Some(&oldest) = self.timestamps.front() {
if now.duration_since(oldest) > self.config.within {
self.timestamps.pop_front();
} else {
break;
}
}
if self.timestamps.len() >= self.config.max_restarts as usize {
return false; }
self.timestamps.push_back(now);
true
}
}
pub type SpawnFn = Box<dyn Fn() -> Pin<Box<dyn Future<Output = SageResult<()>> + Send>> + Send>;
struct ChildHandle {
name: String,
restart_policy: RestartPolicy,
spawn_fn: SpawnFn,
handle: Option<JoinHandle<SageResult<()>>>,
}
impl ChildHandle {
fn new(name: String, restart_policy: RestartPolicy, spawn_fn: SpawnFn) -> Self {
Self {
name,
restart_policy,
spawn_fn,
handle: None,
}
}
fn spawn(&mut self) {
let future = (self.spawn_fn)();
self.handle = Some(tokio::spawn(future));
}
fn is_running(&self) -> bool {
self.handle
.as_ref()
.map(|h| !h.is_finished())
.unwrap_or(false)
}
fn take_handle(&mut self) -> Option<JoinHandle<SageResult<()>>> {
self.handle.take()
}
}
pub struct Supervisor {
strategy: Strategy,
children: Vec<ChildHandle>,
restart_tracker: RestartTracker,
}
impl Supervisor {
pub fn new(strategy: Strategy, config: RestartConfig) -> Self {
Self {
strategy,
children: Vec::new(),
restart_tracker: RestartTracker::new(config),
}
}
pub fn add_child<F, Fut>(
&mut self,
name: impl Into<String>,
restart_policy: RestartPolicy,
spawn_fn: F,
) where
F: Fn() -> Fut + Send + 'static,
Fut: Future<Output = SageResult<()>> + Send + 'static,
{
let spawn_fn: SpawnFn = Box::new(move || Box::pin(spawn_fn()));
self.children
.push(ChildHandle::new(name.into(), restart_policy, spawn_fn));
}
pub async fn run(&mut self) -> SageResult<()> {
for child in &mut self.children {
child.spawn();
}
loop {
let (index, result) = self.wait_for_child_exit().await;
if index.is_none() {
break;
}
let index = index.unwrap();
let child_name = self.children[index].name.clone();
let restart_policy = self.children[index].restart_policy;
let should_restart = match (restart_policy, &result) {
(RestartPolicy::Permanent, _) => true,
(RestartPolicy::Transient, Err(_)) => true,
(RestartPolicy::Transient, Ok(_)) => false,
(RestartPolicy::Temporary, _) => false,
};
if should_restart {
if !self.restart_tracker.record_restart() {
return Err(SageError::Supervisor(format!(
"Maximum restart intensity reached for supervisor (child '{}' failed too many times)",
child_name
)));
}
match self.strategy {
Strategy::OneForOne => {
self.restart_child(index);
}
Strategy::OneForAll => {
self.restart_all();
}
Strategy::RestForOne => {
self.restart_rest(index);
}
}
}
if !self.any_running() {
break;
}
}
Ok(())
}
async fn wait_for_child_exit(&mut self) -> (Option<usize>, SageResult<()>) {
use futures::future::select_all;
let handles_with_indices: Vec<(usize, JoinHandle<SageResult<()>>)> = self
.children
.iter_mut()
.enumerate()
.filter_map(|(i, c)| c.take_handle().map(|h| (i, h)))
.collect();
if handles_with_indices.is_empty() {
return (None, Ok(()));
}
let indices: Vec<usize> = handles_with_indices.iter().map(|(i, _)| *i).collect();
let handles: Vec<JoinHandle<SageResult<()>>> =
handles_with_indices.into_iter().map(|(_, h)| h).collect();
let (join_result, completed_idx, remaining_handles) = select_all(handles).await;
let child_index = indices[completed_idx];
let final_result =
join_result.unwrap_or_else(|e| Err(SageError::Agent(e.to_string())));
let mut remaining_iter = remaining_handles.into_iter();
for (pos, &original_idx) in indices.iter().enumerate() {
if pos != completed_idx {
if let (Some(handle), Some(child)) =
(remaining_iter.next(), self.children.get_mut(original_idx))
{
child.handle = Some(handle);
}
}
}
(Some(child_index), final_result)
}
fn restart_child(&mut self, index: usize) {
if let Some(child) = self.children.get_mut(index) {
child.spawn();
}
}
fn restart_all(&mut self) {
for child in &mut self.children {
if let Some(handle) = child.take_handle() {
handle.abort();
}
}
for child in &mut self.children {
child.spawn();
}
}
fn restart_rest(&mut self, from_index: usize) {
for child in self.children.iter_mut().skip(from_index) {
if let Some(handle) = child.take_handle() {
handle.abort();
}
}
for child in self.children.iter_mut().skip(from_index) {
child.spawn();
}
}
fn any_running(&self) -> bool {
self.children.iter().any(|c| c.is_running())
}
}
}
#[cfg(not(target_arch = "wasm32"))]
pub use native::{SpawnFn, Supervisor};
#[cfg(target_arch = "wasm32")]
mod wasm_stub {
use super::*;
use std::future::Future;
pub struct Supervisor {
_strategy: Strategy,
}
impl Supervisor {
pub fn new(strategy: Strategy, _config: RestartConfig) -> Self {
Self {
_strategy: strategy,
}
}
pub fn add_child<F, Fut>(
&mut self,
_name: impl Into<String>,
_restart_policy: RestartPolicy,
_spawn_fn: F,
) where
F: Fn() -> Fut + 'static,
Fut: Future<Output = SageResult<()>> + 'static,
{
}
pub async fn run(&mut self) -> SageResult<()> {
Err(SageError::Supervisor(
"Supervision trees are not yet supported in the WASM target".to_string(),
))
}
}
}
#[cfg(target_arch = "wasm32")]
pub use wasm_stub::Supervisor;
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
#[tokio::test]
async fn test_one_for_one_restart() {
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let mut supervisor = Supervisor::new(Strategy::OneForOne, RestartConfig::default());
supervisor.add_child("Worker", RestartPolicy::Transient, move || {
let counter = counter_clone.clone();
async move {
let count = counter.fetch_add(1, Ordering::SeqCst);
if count < 2 {
Err(SageError::Agent("Simulated failure".to_string()))
} else {
Ok(())
}
}
});
let result = supervisor.run().await;
assert!(result.is_ok(), "supervisor failed: {:?}", result);
assert_eq!(counter.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn test_transient_no_restart_on_success() {
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let mut supervisor = Supervisor::new(Strategy::OneForOne, RestartConfig::default());
supervisor.add_child("Worker", RestartPolicy::Transient, move || {
let counter = counter_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
Ok(())
}
});
let result = supervisor.run().await;
assert!(result.is_ok());
assert_eq!(counter.load(Ordering::SeqCst), 1); }
#[tokio::test]
async fn test_temporary_never_restarts() {
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let mut supervisor = Supervisor::new(Strategy::OneForOne, RestartConfig::default());
supervisor.add_child("Worker", RestartPolicy::Temporary, move || {
let counter = counter_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
Err(SageError::Agent("Simulated failure".to_string()))
}
});
let result = supervisor.run().await;
assert!(result.is_ok()); assert_eq!(counter.load(Ordering::SeqCst), 1); }
#[tokio::test]
async fn test_circuit_breaker() {
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let config = RestartConfig {
max_restarts: 3,
within: Duration::from_secs(60),
};
let mut supervisor = Supervisor::new(Strategy::OneForOne, config);
supervisor.add_child("Worker", RestartPolicy::Permanent, move || {
let counter = counter_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
Err(SageError::Agent("Always fails".to_string()))
}
});
let result = supervisor.run().await;
assert!(result.is_err()); assert!(counter.load(Ordering::SeqCst) <= 4); }
#[tokio::test]
async fn test_permanent_restarts_on_success() {
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let config = RestartConfig {
max_restarts: 3,
within: Duration::from_secs(60),
};
let mut supervisor = Supervisor::new(Strategy::OneForOne, config);
supervisor.add_child("Worker", RestartPolicy::Permanent, move || {
let counter = counter_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
Ok(()) }
});
let result = supervisor.run().await;
assert!(result.is_err());
assert!(counter.load(Ordering::SeqCst) <= 4);
}
#[tokio::test]
async fn test_rest_for_one_restarts_downstream() {
let counter1 = Arc::new(AtomicU32::new(0));
let counter2 = Arc::new(AtomicU32::new(0));
let counter3 = Arc::new(AtomicU32::new(0));
let counter1_clone = counter1.clone();
let counter2_clone = counter2.clone();
let counter3_clone = counter3.clone();
let mut supervisor = Supervisor::new(Strategy::RestForOne, RestartConfig::default());
supervisor.add_child("Child1", RestartPolicy::Temporary, move || {
let counter = counter1_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(())
}
});
supervisor.add_child("Child2", RestartPolicy::Transient, move || {
let counter = counter2_clone.clone();
async move {
let count = counter.fetch_add(1, Ordering::SeqCst);
if count < 2 {
Err(SageError::Agent("Simulated failure".to_string()))
} else {
Ok(())
}
}
});
supervisor.add_child("Child3", RestartPolicy::Temporary, move || {
let counter = counter3_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(())
}
});
let result = supervisor.run().await;
assert!(result.is_ok(), "supervisor failed: {:?}", result);
assert_eq!(
counter1.load(Ordering::SeqCst),
1,
"Child1 should run only once"
);
assert_eq!(
counter2.load(Ordering::SeqCst),
3,
"Child2 should run 3 times"
);
assert!(
counter3.load(Ordering::SeqCst) >= 2,
"Child3 should be restarted at least once with RestForOne, got {}",
counter3.load(Ordering::SeqCst)
);
}
#[tokio::test]
async fn test_one_for_all_restarts_all() {
let counter1 = Arc::new(AtomicU32::new(0));
let counter2 = Arc::new(AtomicU32::new(0));
let counter1_clone = counter1.clone();
let counter2_clone = counter2.clone();
let mut supervisor = Supervisor::new(Strategy::OneForAll, RestartConfig::default());
supervisor.add_child("Child1", RestartPolicy::Temporary, move || {
let counter = counter1_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(100)).await;
Ok(())
}
});
supervisor.add_child("Child2", RestartPolicy::Transient, move || {
let counter = counter2_clone.clone();
async move {
let count = counter.fetch_add(1, Ordering::SeqCst);
if count < 2 {
Err(SageError::Agent("Simulated failure".to_string()))
} else {
tokio::time::sleep(Duration::from_millis(10)).await;
Ok(())
}
}
});
let result = supervisor.run().await;
assert!(result.is_ok(), "supervisor failed: {:?}", result);
assert_eq!(
counter2.load(Ordering::SeqCst),
3,
"Child2 should run 3 times"
);
assert!(
counter1.load(Ordering::SeqCst) >= 2,
"Child1 should be restarted at least once with OneForAll, got {}",
counter1.load(Ordering::SeqCst)
);
}
}