use std::collections::BTreeMap;
use std::sync::{Arc, Mutex, OnceLock};
use aion_core::ClusterEvent;
use aion_store::{
DesiredState, WorkerArtifactRef, WorkerDeployment, WorkerDeploymentListing,
WorkerDeploymentStore,
};
use crate::cluster_publisher::ClusterEventPublisher;
use super::error::SupervisionError;
use super::executable::ManagedExecutable;
use super::instance::{InstanceConfig, InstanceHandle, InstanceSnapshot};
use super::policy::{SupervisionPolicy, UNCOMMISSIONED_REMEDY};
use super::status::{ManagedWorkerReport, ManagedWorkerState, ManagedWorkerStatus};
#[derive(Clone, Debug)]
struct Commission {
policy: SupervisionPolicy,
executable: ManagedExecutable,
}
pub struct WorkerSupervisor {
store: Arc<dyn WorkerDeploymentStore>,
commission: OnceLock<Commission>,
instances: Mutex<BTreeMap<String, InstanceHandle>>,
publisher: ClusterEventPublisher,
}
impl WorkerSupervisor {
#[must_use]
pub fn new(store: Arc<dyn WorkerDeploymentStore>, publisher: ClusterEventPublisher) -> Self {
Self {
store,
commission: OnceLock::new(),
instances: Mutex::new(BTreeMap::new()),
publisher,
}
}
pub fn commission(&self, policy: SupervisionPolicy, executable: ManagedExecutable) -> bool {
self.commission
.set(Commission { policy, executable })
.is_ok()
}
#[must_use]
pub fn policy(&self) -> Option<SupervisionPolicy> {
self.commission.get().map(|commission| commission.policy)
}
#[must_use]
pub fn is_commissioned(&self) -> bool {
self.commission.get().is_some()
}
fn require_commission(&self) -> Result<&Commission, SupervisionError> {
self.commission
.get()
.ok_or(SupervisionError::NotCommissioned)
}
fn instances(
&self,
) -> Result<std::sync::MutexGuard<'_, BTreeMap<String, InstanceHandle>>, SupervisionError> {
self.instances
.lock()
.map_err(|poison| SupervisionError::StatePoisoned {
detail: poison.to_string(),
})
}
async fn record(&self, name: &str) -> Result<WorkerDeployment, SupervisionError> {
self.store
.get_worker_deployment(name)
.await
.map_err(|source| SupervisionError::Store { source })?
.ok_or_else(|| SupervisionError::UnknownDeployment {
name: name.to_owned(),
})
}
async fn listing(&self) -> Result<WorkerDeploymentListing, SupervisionError> {
self.store
.list_worker_deployments()
.await
.map_err(|source| SupervisionError::Store { source })
}
pub async fn start(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
let commission = self.require_commission()?.clone();
let record = self.record(name).await?;
let record = if record.desired == DesiredState::Running {
record
} else {
let record = self
.store
.set_desired_state(name, DesiredState::Running)
.await
.map_err(|source| SupervisionError::Store { source })?
.ok_or_else(|| SupervisionError::UnknownDeployment {
name: name.to_owned(),
})?;
self.publish_desired_state(&record);
record
};
let snapshot = {
let mut instances = self.instances()?;
let live = instances
.get(name)
.is_some_and(|handle| !handle.is_finished());
if !live {
let handle = InstanceHandle::start(self.instance_config(&record, &commission));
drop(instances.insert(record.name.clone(), handle));
}
instances
.get(name)
.map(InstanceHandle::snapshot)
.transpose()?
};
Ok(status_of(&record, snapshot, self.is_commissioned()))
}
pub async fn stop(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
let record = self
.store
.set_desired_state(name, DesiredState::Stopped)
.await
.map_err(|source| SupervisionError::Store { source })?
.ok_or_else(|| SupervisionError::UnknownDeployment {
name: name.to_owned(),
})?;
self.publish_desired_state(&record);
let handle = self.instances()?.remove(name);
if let Some(handle) = handle {
handle.stop(name).await?;
}
Ok(status_of(&record, None, self.is_commissioned()))
}
pub async fn restart(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
self.require_commission()?;
self.record(name).await?;
let handle = self.instances()?.remove(name);
if let Some(handle) = handle {
handle.stop(name).await?;
}
self.start(name).await
}
pub async fn reconcile(&self) -> Result<usize, SupervisionError> {
let commission = self.require_commission()?.clone();
let listing = self.listing().await?;
let mut supervised = 0_usize;
for record in listing
.deployments
.iter()
.filter(|record| record.desired == DesiredState::Running)
{
let mut instances = self.instances()?;
let live = instances
.get(&record.name)
.is_some_and(|handle| !handle.is_finished());
if !live {
let handle = InstanceHandle::start(self.instance_config(record, &commission));
drop(instances.insert(record.name.clone(), handle));
}
supervised = supervised.saturating_add(1);
}
for poisoned in &listing.undecodable {
tracing::error!(
worker = poisoned.name.as_str(),
error = poisoned.error.as_str(),
"worker deployment record could not be decoded; it is not being supervised"
);
}
Ok(supervised)
}
pub async fn report(&self) -> Result<ManagedWorkerReport, SupervisionError> {
let listing = self.listing().await?;
let commissioned = self.is_commissioned();
let mut workers = Vec::with_capacity(listing.deployments.len());
{
let instances = self.instances()?;
for record in &listing.deployments {
let snapshot = instances
.get(&record.name)
.map(InstanceHandle::snapshot)
.transpose()?;
workers.push(status_of(record, snapshot, commissioned));
}
}
Ok(ManagedWorkerReport {
commissioned,
remedy: (!commissioned).then(|| UNCOMMISSIONED_REMEDY.to_owned()),
workers,
undecodable: listing
.undecodable
.iter()
.map(|poisoned| poisoned.name.clone())
.collect(),
})
}
pub async fn shutdown(&self) -> Vec<SupervisionError> {
let handles = match self.instances() {
Ok(mut instances) => std::mem::take(&mut *instances),
Err(error) => return vec![error],
};
let mut failures = Vec::new();
for (name, handle) in handles {
if let Err(error) = handle.stop(&name).await {
failures.push(error);
}
}
failures
}
fn publish_desired_state(&self, record: &WorkerDeployment) {
let name = record.name.clone();
let desired_state = record.desired;
drop(
self.publisher
.emit(|meta| ClusterEvent::WorkerDeploymentDesiredStateChanged {
meta,
name,
desired_state,
}),
);
}
fn instance_config(
&self,
record: &WorkerDeployment,
commission: &Commission,
) -> InstanceConfig {
let WorkerArtifactRef::Builtin { verb } = &record.artifact;
InstanceConfig {
name: record.name.clone(),
verb: verb.clone(),
executable: commission.executable.clone(),
policy: commission.policy,
store: Arc::clone(&self.store),
}
}
}
fn status_of(
record: &WorkerDeployment,
snapshot: Option<InstanceSnapshot>,
commissioned: bool,
) -> ManagedWorkerStatus {
let supervised = snapshot.is_some();
let snapshot = snapshot.unwrap_or_default();
let state = snapshot.state.unwrap_or({
if supervised {
ManagedWorkerState::Starting
} else if record.desired == DesiredState::Running && !commissioned {
ManagedWorkerState::Uncommissioned
} else {
ManagedWorkerState::Stopped
}
});
ManagedWorkerStatus {
name: record.name.clone(),
task_queue: record.task_queue.clone(),
desired: record.desired,
state,
pid: snapshot.pid,
process_group: snapshot.process_group,
restarts: snapshot.restarts,
last_exit: snapshot.last_exit,
last_error: snapshot.last_error,
deployed_binary: record.binary.clone(),
spawn_binary: snapshot.spawn_binary,
}
}