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 crate::worker::auto_provision::AutoWorkerOutcome;
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,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Convergence {
Idempotent,
Replacing,
}
#[derive(Debug, Default)]
pub struct FleetShutdownReport {
pub stopped: Vec<String>,
pub failures: Vec<SupervisionError>,
}
pub struct WorkerSupervisor {
store: Arc<dyn WorkerDeploymentStore>,
commission: OnceLock<Commission>,
instances: Mutex<BTreeMap<String, InstanceHandle>>,
publisher: ClusterEventPublisher,
auto_provision: Mutex<BTreeMap<String, AutoWorkerOutcome>>,
gates: Mutex<BTreeMap<String, Arc<tokio::sync::Mutex<()>>>>,
}
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,
auto_provision: Mutex::new(BTreeMap::new()),
gates: Mutex::new(BTreeMap::new()),
}
}
fn gate(&self, name: &str) -> Result<Arc<tokio::sync::Mutex<()>>, SupervisionError> {
let mut gates = self
.gates
.lock()
.map_err(|poison| SupervisionError::StatePoisoned {
detail: poison.to_string(),
})?;
Ok(Arc::clone(
gates
.entry(name.to_owned())
.or_insert_with(|| Arc::new(tokio::sync::Mutex::new(()))),
))
}
pub fn record_auto_provision(&self, outcomes: &[AutoWorkerOutcome]) {
match self.auto_provision.lock() {
Ok(mut log) => {
for outcome in outcomes {
drop(log.insert(outcome.task_queue.clone(), outcome.clone()));
}
}
Err(poison) => tracing::error!(
detail = %poison,
"the auto-provision log is poisoned; this server's managed-worker report will \
not say why a queue has no built-in agent worker"
),
}
}
fn auto_provision_log(&self) -> Vec<AutoWorkerOutcome> {
self.auto_provision
.lock()
.map(|log| log.values().cloned().collect())
.unwrap_or_default()
}
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 gate = self.gate(name)?;
let held = gate.lock().await;
let started = self.start_gated(name).await;
drop(held);
started
}
async fn start_gated(&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 = self.ensure_instance(&record, &commission)?;
Ok(status_of(&record, snapshot, self.is_commissioned()))
}
pub async fn converge(
&self,
name: &str,
mode: Convergence,
) -> Result<ManagedWorkerStatus, SupervisionError> {
let gate = self.gate(name)?;
let held = gate.lock().await;
let converged = self.converge_gated(name, mode).await;
drop(held);
converged
}
async fn converge_gated(
&self,
name: &str,
mode: Convergence,
) -> Result<ManagedWorkerStatus, SupervisionError> {
let record = self.record(name).await?;
if record.desired == DesiredState::Stopped {
self.tear_down(name).await?;
return Ok(status_of(&record, None, self.is_commissioned()));
}
let commission = self.require_commission()?.clone();
if mode == Convergence::Replacing {
self.tear_down(name).await?;
}
let snapshot = self.ensure_instance(&record, &commission)?;
Ok(status_of(&record, snapshot, self.is_commissioned()))
}
pub async fn forget(&self, name: &str) -> Result<(), SupervisionError> {
let gate = self.gate(name)?;
let held = gate.lock().await;
let forgotten = self.tear_down(name).await;
drop(held);
forgotten
}
async fn tear_down(&self, name: &str) -> Result<(), SupervisionError> {
let handle = self.instances()?.remove(name);
let Some(handle) = handle else {
return Ok(());
};
let identity = handle.snapshot().ok().map(|snapshot| ProcessIdentity {
pid: snapshot.pid,
process_group: snapshot.process_group,
});
handle
.stop(name)
.await
.map_err(|error| attach_process_identity(error, identity))
}
fn ensure_instance(
&self,
record: &WorkerDeployment,
commission: &Commission,
) -> Result<Option<InstanceSnapshot>, SupervisionError> {
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));
}
instances
.get(&record.name)
.map(InstanceHandle::snapshot)
.transpose()
}
pub async fn stop(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
let gate = self.gate(name)?;
let held = gate.lock().await;
let stopped = self.stop_gated(name).await;
drop(held);
stopped
}
async fn stop_gated(&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);
self.tear_down(name).await?;
Ok(status_of(&record, None, self.is_commissioned()))
}
pub async fn restart(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
let gate = self.gate(name)?;
let held = gate.lock().await;
let restarted = self.restart_gated(name).await;
drop(held);
restarted
}
async fn restart_gated(&self, name: &str) -> Result<ManagedWorkerStatus, SupervisionError> {
self.require_commission()?;
self.record(name).await?;
self.tear_down(name).await?;
self.start_gated(name).await
}
pub async fn reconcile(&self) -> Result<usize, SupervisionError> {
drop(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)
{
match self.converge(&record.name, Convergence::Idempotent).await {
Ok(_) => supervised = supervised.saturating_add(1),
Err(error) => tracing::error!(
worker = record.name.as_str(),
%error,
"worker deployment could not be converged; it is not being supervised"
),
}
}
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(),
auto_provision: self.auto_provision_log(),
})
}
pub async fn shutdown(&self) -> FleetShutdownReport {
let handles = match self.instances() {
Ok(mut instances) => std::mem::take(&mut *instances),
Err(error) => {
return FleetShutdownReport {
stopped: Vec::new(),
failures: vec![error],
};
}
};
let mut report = FleetShutdownReport {
stopped: Vec::new(),
failures: Vec::new(),
};
for (name, handle) in handles {
match handle.stop(&name).await {
Ok(()) => report.stopped.push(name),
Err(error) => report.failures.push(error),
}
}
report
}
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),
}
}
}
struct ProcessIdentity {
pid: Option<u32>,
process_group: Option<i32>,
}
fn attach_process_identity(
error: SupervisionError,
identity: Option<ProcessIdentity>,
) -> SupervisionError {
match (error, identity) {
(SupervisionError::StopIncomplete { name, detail }, Some(identity)) => {
let pid = identity
.pid
.map_or_else(|| String::from("unknown"), |pid| pid.to_string());
let group = identity
.process_group
.map_or_else(|| String::from("unknown"), |group| group.to_string());
SupervisionError::StopIncomplete {
name,
detail: format!("{detail} (last observed pid {pid}, process group {group})"),
}
}
(error, _) => error,
}
}
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,
}
}