use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use anyhow::Result;
use async_trait::async_trait;
use tokio::sync::{mpsc, oneshot, watch};
use workload_spec::{BackoffPolicy, RestartPolicy};
use crate::WorkloadStatus;
const ALWAYS_RESTART_DELAY: Duration = Duration::from_secs(1);
pub(crate) type Exit = std::io::Result<std::process::ExitStatus>;
pub(crate) struct Completion {
pub succeeded: bool,
pub exit_code: i32,
pub terminal: WorkloadStatus,
}
pub(crate) enum Ctrl<I> {
Teardown(oneshot::Sender<()>),
Restart(oneshot::Sender<Result<u32>>),
#[cfg_attr(not(feature = "native-integration"), allow(dead_code))]
Adopt {
replacement: I,
ack: oneshot::Sender<()>,
},
}
#[async_trait]
pub(crate) trait Supervised: Send + Sync + 'static {
type Instance: Send + 'static;
fn pid(&self, inst: &Self::Instance) -> u32;
async fn wait(&self, inst: &mut Self::Instance) -> Exit;
async fn start(&self) -> Result<Self::Instance>;
async fn stop(&self, inst: &mut Self::Instance);
async fn settle(&self, exit: Exit) -> Completion;
fn discard(&self, _replaced: Self::Instance) {}
}
enum Ev<I> {
Exited(Exit),
Ctrl(Option<Ctrl<I>>),
}
fn backoff_delay(b: &BackoffPolicy, attempt: u32) -> Duration {
let factor = (b.multiplier as f64).powi(attempt.saturating_sub(1) as i32);
let ms = (b.initial_ms as f64 * factor).min(b.max_ms as f64);
Duration::from_millis(ms as u64)
}
fn now_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
async fn park<S: Supervised>(
sup: &S,
ctrl_rx: &mut mpsc::Receiver<Ctrl<S::Instance>>,
pid: &Arc<AtomicU32>,
status_tx: &watch::Sender<WorkloadStatus>,
) -> Option<S::Instance> {
loop {
match ctrl_rx.recv().await {
None => return None,
Some(Ctrl::Teardown(ack)) => {
let _ = status_tx.send(WorkloadStatus::Stopped);
let _ = ack.send(());
return None;
}
Some(Ctrl::Adopt { replacement, ack }) => {
pid.store(sup.pid(&replacement), Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
let _ = ack.send(());
return Some(replacement);
}
Some(Ctrl::Restart(ack)) => match sup.start().await {
Ok(new) => {
let p = sup.pid(&new);
pid.store(p, Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
let _ = ack.send(Ok(p));
return Some(new);
}
Err(e) => {
let _ = status_tx.send(WorkloadStatus::Failed {
reason: format!("restart respawn failed: {e}"),
oom_killed: false,
});
let _ = ack.send(Err(e));
}
},
}
}
}
pub(crate) async fn supervise<S: Supervised>(
sup: S,
policy: RestartPolicy,
initial: S::Instance,
pid: Arc<AtomicU32>,
status_tx: watch::Sender<WorkloadStatus>,
mut ctrl_rx: mpsc::Receiver<Ctrl<S::Instance>>,
) {
let mut inst = initial;
let mut failure_streak: u32 = 0;
let mut restart_count: u32 = 0;
loop {
let ev = tokio::select! {
r = sup.wait(&mut inst) => Ev::Exited(r),
c = ctrl_rx.recv() => Ev::Ctrl(c),
};
match ev {
Ev::Ctrl(None) => {
sup.stop(&mut inst).await;
return;
}
Ev::Ctrl(Some(Ctrl::Teardown(ack))) => {
sup.stop(&mut inst).await;
pid.store(0, Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Stopped);
let _ = ack.send(());
return;
}
Ev::Ctrl(Some(Ctrl::Restart(ack))) => {
sup.stop(&mut inst).await;
pid.store(0, Ordering::SeqCst);
match sup.start().await {
Ok(new) => {
restart_count += 1;
failure_streak = 0;
let p = sup.pid(&new);
pid.store(p, Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
inst = new;
let _ = ack.send(Ok(p));
}
Err(e) => {
let _ = status_tx.send(WorkloadStatus::Failed {
reason: format!("restart respawn failed: {e}"),
oom_killed: false,
});
let _ = ack.send(Err(e));
match park(&sup, &mut ctrl_rx, &pid, &status_tx).await {
Some(new) => {
restart_count += 1;
failure_streak = 0;
inst = new;
}
None => return,
}
}
}
}
Ev::Ctrl(Some(Ctrl::Adopt { replacement, ack })) => {
sup.discard(inst);
restart_count += 1;
failure_streak = 0;
pid.store(sup.pid(&replacement), Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
inst = replacement;
let _ = ack.send(());
}
Ev::Exited(exit) => {
pid.store(0, Ordering::SeqCst);
let done = sup.settle(exit).await;
let should_restart = match &policy {
RestartPolicy::Always => true,
RestartPolicy::Never => false,
RestartPolicy::OnFailure { max_attempts, .. } => {
!done.succeeded && failure_streak < *max_attempts
}
};
if !should_restart {
let _ = status_tx.send(done.terminal);
match park(&sup, &mut ctrl_rx, &pid, &status_tx).await {
Some(new) => {
restart_count += 1;
failure_streak = 0;
inst = new;
continue;
}
None => return,
}
}
restart_count += 1;
if !done.succeeded {
failure_streak += 1;
}
let _ = status_tx.send(WorkloadStatus::Restarting {
last_exit_code: done.exit_code,
restart_count,
last_finished_at_unix_ms: now_ms(),
});
let delay = match &policy {
RestartPolicy::Always => ALWAYS_RESTART_DELAY,
RestartPolicy::OnFailure { backoff, .. } => {
backoff_delay(backoff, failure_streak)
}
RestartPolicy::Never => Duration::ZERO,
};
let interrupted = tokio::select! {
_ = tokio::time::sleep(delay) => None,
c = ctrl_rx.recv() => Some(c),
};
match interrupted {
None => {} Some(None) => return,
Some(Some(Ctrl::Teardown(ack))) => {
let _ = status_tx.send(WorkloadStatus::Stopped);
let _ = ack.send(());
return;
}
Some(Some(Ctrl::Adopt { replacement, ack })) => {
failure_streak = 0;
pid.store(sup.pid(&replacement), Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
inst = replacement;
let _ = ack.send(());
continue;
}
Some(Some(Ctrl::Restart(ack))) => match sup.start().await {
Ok(new) => {
failure_streak = 0;
let p = sup.pid(&new);
pid.store(p, Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
inst = new;
let _ = ack.send(Ok(p));
continue;
}
Err(e) => {
let _ = status_tx.send(WorkloadStatus::Failed {
reason: format!("restart respawn failed: {e}"),
oom_killed: false,
});
let _ = ack.send(Err(e));
match park(&sup, &mut ctrl_rx, &pid, &status_tx).await {
Some(new) => {
failure_streak = 0;
inst = new;
continue;
}
None => return,
}
}
},
}
match sup.start().await {
Ok(new) => {
pid.store(sup.pid(&new), Ordering::SeqCst);
let _ = status_tx.send(WorkloadStatus::Running);
inst = new;
}
Err(e) => {
let _ = status_tx.send(WorkloadStatus::Failed {
reason: format!("respawn failed: {e}"),
oom_killed: false,
});
match park(&sup, &mut ctrl_rx, &pid, &status_tx).await {
Some(new) => {
failure_streak = 0;
inst = new;
}
None => return,
}
}
}
}
}
}
}