use crate::{ManagedService, ServiceRuntimeState, SupervisorError, SupervisorResult};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::mpsc::{self, Receiver, RecvTimeoutError, SyncSender, TrySendError};
use std::sync::{Arc, Mutex};
use std::thread::JoinHandle;
use std::time::{Duration, Instant};
pub(crate) const DEFAULT_RESTART_QUEUE_CAPACITY: usize = 64;
pub(crate) const DEFAULT_RESTART_WORKERS: usize = 2;
const RESTART_THREAD_STACK_BYTES: usize = 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RestartState {
None,
Scheduled {
execute_at_ms: u64,
},
Stopping,
Starting,
Backoff,
Failed,
}
pub(crate) struct RestartCommand {
pub service: Arc<dyn ManagedService>,
pub attempt: u64,
}
pub(crate) enum RestartOutcome {
Restarted,
Orphaned,
Failed,
Cancelled,
}
pub(crate) struct RestartCompletion {
pub service_id: String,
pub attempt: u64,
pub outcome: RestartOutcome,
}
pub(crate) struct RestartExecutor {
sender: Mutex<Option<SyncSender<RestartCommand>>>,
commands: Arc<Mutex<Receiver<RestartCommand>>>,
completions: Mutex<Receiver<RestartCompletion>>,
workers: Mutex<Vec<JoinHandle<()>>>,
cancellation: Arc<AtomicBool>,
healthy: Arc<AtomicBool>,
pending: Arc<AtomicU64>,
queue_capacity: usize,
worker_count: usize,
}
impl RestartExecutor {
pub fn new(queue_capacity: usize, worker_count: usize) -> Self {
let queue_capacity = queue_capacity.max(1);
let worker_count = worker_count.max(1);
let (sender, receiver) = mpsc::sync_channel::<RestartCommand>(queue_capacity);
let completion_capacity = queue_capacity.saturating_add(worker_count);
let (completion_sender, completions) =
mpsc::sync_channel::<RestartCompletion>(completion_capacity);
let receiver = Arc::new(Mutex::new(receiver));
let cancellation = Arc::new(AtomicBool::new(false));
let healthy = Arc::new(AtomicBool::new(true));
let pending = Arc::new(AtomicU64::new(0));
let workers = spawn_workers(
worker_count,
Arc::clone(&receiver),
completion_sender,
Arc::clone(&cancellation),
Arc::clone(&healthy),
Arc::clone(&pending),
);
Self {
sender: Mutex::new(Some(sender)),
commands: receiver,
completions: Mutex::new(completions),
workers: Mutex::new(workers),
cancellation,
healthy,
pending,
queue_capacity,
worker_count,
}
}
pub fn schedule(&self, command: RestartCommand) -> SupervisorResult<()> {
if self.cancellation.load(Ordering::Acquire) {
return Err(SupervisorError::RestartExecutorStopped);
}
let sender = self
.sender
.lock()
.map_err(|_| SupervisorError::RestartExecutorStopped)?;
let sender = sender
.as_ref()
.ok_or(SupervisorError::RestartExecutorStopped)?;
self.pending.fetch_add(1, Ordering::AcqRel);
match sender.try_send(command) {
Ok(()) => Ok(()),
Err(TrySendError::Full(_)) => {
self.pending.fetch_sub(1, Ordering::AcqRel);
Err(SupervisorError::RestartQueueFull)
}
Err(TrySendError::Disconnected(_)) => {
self.pending.fetch_sub(1, Ordering::AcqRel);
Err(SupervisorError::RestartExecutorStopped)
}
}
}
pub fn drain_completions(&self) -> Vec<RestartCompletion> {
let Ok(receiver) = self.completions.lock() else {
self.healthy.store(false, Ordering::Release);
return Vec::new();
};
receiver.try_iter().collect()
}
pub fn snapshot(&self) -> crate::RestartExecutorSnapshot {
let workers_healthy = self
.workers
.lock()
.map(|workers| workers.iter().all(|worker| !worker.is_finished()))
.unwrap_or(false);
crate::RestartExecutorSnapshot {
healthy: self.healthy.load(Ordering::Acquire)
&& workers_healthy
&& !self.cancellation.load(Ordering::Acquire),
pending: self.pending.load(Ordering::Acquire),
queue_capacity: self.queue_capacity,
worker_count: self.worker_count,
}
}
pub fn shutdown(&self, timeout: Duration) -> bool {
self.cancellation.store(true, Ordering::Release);
if let Ok(mut sender) = self.sender.lock() {
sender.take();
} else {
self.healthy.store(false, Ordering::Release);
}
let deadline = Instant::now().checked_add(timeout);
while deadline.is_none_or(|deadline| Instant::now() < deadline) {
let complete = self
.workers
.lock()
.map(|workers| workers.iter().all(JoinHandle::is_finished))
.unwrap_or(false);
if complete {
break;
}
std::thread::sleep(Duration::from_millis(5));
}
let Ok(mut workers) = self.workers.lock() else {
self.healthy.store(false, Ordering::Release);
return false;
};
let all_finished = workers.iter().all(JoinHandle::is_finished);
for worker in workers.drain(..).filter(JoinHandle::is_finished) {
if worker.join().is_err() {
self.healthy.store(false, Ordering::Release);
}
}
drain_retained(&self.commands, &self.healthy);
drain_retained(&self.completions, &self.healthy);
self.pending.store(0, Ordering::Release);
self.healthy.store(false, Ordering::Release);
all_finished
}
}
fn spawn_workers(
count: usize,
receiver: Arc<Mutex<Receiver<RestartCommand>>>,
completions: SyncSender<RestartCompletion>,
cancellation: Arc<AtomicBool>,
healthy: Arc<AtomicBool>,
pending: Arc<AtomicU64>,
) -> Vec<JoinHandle<()>> {
(0..count)
.filter_map(|index| {
let receiver = Arc::clone(&receiver);
let completions = completions.clone();
let cancellation = Arc::clone(&cancellation);
let healthy = Arc::clone(&healthy);
let worker_health = Arc::clone(&healthy);
let pending = Arc::clone(&pending);
std::thread::Builder::new()
.name(format!("appcore-restart-{index}"))
.stack_size(RESTART_THREAD_STACK_BYTES)
.spawn(move || {
restart_worker(receiver, completions, cancellation, worker_health, pending)
})
.map_err(|_| healthy.store(false, Ordering::Release))
.ok()
})
.collect()
}
fn restart_worker(
receiver: Arc<Mutex<Receiver<RestartCommand>>>,
completions: SyncSender<RestartCompletion>,
cancellation: Arc<AtomicBool>,
healthy: Arc<AtomicBool>,
pending: Arc<AtomicU64>,
) {
loop {
if cancellation.load(Ordering::Acquire) {
return;
}
let command = match receiver
.lock()
.map(|receiver| receiver.recv_timeout(Duration::from_millis(25)))
{
Ok(Ok(command)) => command,
Ok(Err(RecvTimeoutError::Timeout)) => continue,
Ok(Err(RecvTimeoutError::Disconnected)) | Err(_) => return,
};
let service_id = command.service.descriptor().name().to_string();
let attempt = command.attempt;
let completion = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
execute_restart(command, &cancellation)
}))
.unwrap_or(RestartCompletion {
service_id,
attempt,
outcome: RestartOutcome::Failed,
});
send_completion(&completions, completion, &cancellation, &healthy);
decrement_pending(&pending);
}
}
fn send_completion(
completions: &SyncSender<RestartCompletion>,
mut completion: RestartCompletion,
cancellation: &AtomicBool,
healthy: &AtomicBool,
) {
loop {
if cancellation.load(Ordering::Acquire) {
return;
}
match completions.try_send(completion) {
Ok(()) => return,
Err(TrySendError::Full(retained)) => {
completion = retained;
std::thread::sleep(Duration::from_millis(1));
}
Err(TrySendError::Disconnected(_)) => {
healthy.store(false, Ordering::Release);
return;
}
}
}
}
fn drain_retained<T>(receiver: &Mutex<Receiver<T>>, healthy: &AtomicBool) {
match receiver.lock() {
Ok(receiver) => receiver.try_iter().for_each(drop),
Err(_) => healthy.store(false, Ordering::Release),
}
}
fn decrement_pending(pending: &AtomicU64) {
let _ = pending.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| {
Some(value.saturating_sub(1))
});
}
fn execute_restart(command: RestartCommand, cancellation: &AtomicBool) -> RestartCompletion {
let service_id = command.service.descriptor().name().to_string();
let timeout = command
.service
.descriptor()
.restart_policy()
.shutdown_timeout;
let outcome = match command.service.stop(timeout) {
Err(_) if command.service.runtime_state() == ServiceRuntimeState::Orphaned => {
RestartOutcome::Orphaned
}
Err(_) => RestartOutcome::Failed,
Ok(()) if cancellation.load(Ordering::Acquire) => RestartOutcome::Cancelled,
Ok(()) => match command.service.start() {
Ok(()) => RestartOutcome::Restarted,
Err(_) => RestartOutcome::Failed,
},
};
RestartCompletion {
service_id,
attempt: command.attempt,
outcome,
}
}
#[cfg(test)]
#[path = "restart_executor_tests.rs"]
mod tests;