use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use aion_core::{RunId, WorkflowId};
use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
use tokio::task::JoinHandle;
use crate::EngineError;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct CompletionRetryKey {
pub(crate) workflow_id: WorkflowId,
pub(crate) run_id: RunId,
pub(crate) monitor_pid: super::Pid,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ArmOutcome {
Armed,
AlreadyArmed,
EpochClosed,
}
pub(crate) struct EngineTaskRuntime {
runtime: Mutex<Option<tokio::runtime::Runtime>>,
watches: DashMap<(u64, WorkflowId), JoinHandle<()>>,
spawn_retries: DashMap<WorkflowId, JoinHandle<()>>,
completion_retries: DashMap<CompletionRetryKey, JoinHandle<()>>,
shutting_down: AtomicBool,
}
impl EngineTaskRuntime {
pub(crate) fn new() -> Result<Self, EngineError> {
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.thread_name("aion-engine-tasks")
.enable_all()
.build()
.map_err(|error| EngineError::Runtime {
reason: format!("failed to start the engine-task runtime: {error}"),
})?;
Ok(Self {
runtime: Mutex::new(Some(runtime)),
watches: DashMap::new(),
spawn_retries: DashMap::new(),
completion_retries: DashMap::new(),
shutting_down: AtomicBool::new(false),
})
}
pub(crate) fn arm_watch<F>(&self, parent_pid: u64, child_id: WorkflowId, task: F) -> ArmOutcome
where
F: Future<Output = ()> + Send + 'static,
{
Self::arm(
&self.shutting_down,
&self.runtime,
&self.watches,
(parent_pid, child_id),
task,
)
}
pub(crate) fn arm_spawn_retry<F>(&self, child_id: WorkflowId, task: F) -> ArmOutcome
where
F: Future<Output = ()> + Send + 'static,
{
Self::arm(
&self.shutting_down,
&self.runtime,
&self.spawn_retries,
child_id,
task,
)
}
pub(crate) fn arm_completion_retry<F>(&self, lease: CompletionRetryKey, task: F) -> ArmOutcome
where
F: Future<Output = ()> + Send + 'static,
{
Self::arm(
&self.shutting_down,
&self.runtime,
&self.completion_retries,
lease,
task,
)
}
fn arm<K, F>(
shutting_down: &AtomicBool,
runtime: &Mutex<Option<tokio::runtime::Runtime>>,
registry: &DashMap<K, JoinHandle<()>>,
key: K,
task: F,
) -> ArmOutcome
where
K: std::hash::Hash + Eq + Clone,
F: Future<Output = ()> + Send + 'static,
{
if shutting_down.load(Ordering::Acquire) {
return ArmOutcome::EpochClosed;
}
let handle = {
let guard = match runtime.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
let Some(owned) = guard.as_ref() else {
return ArmOutcome::EpochClosed;
};
owned.handle().clone()
};
let undo_key = key.clone();
let spawned_id;
match registry.entry(key) {
Entry::Occupied(slot) => {
if slot.get().is_finished() {
let spawned = handle.spawn(task);
spawned_id = spawned.id();
let (key, _finished) = slot.replace_entry(spawned);
let _ = key;
} else {
return ArmOutcome::AlreadyArmed;
}
}
Entry::Vacant(slot) => {
let spawned = handle.spawn(task);
spawned_id = spawned.id();
slot.insert(spawned);
}
}
if shutting_down.load(Ordering::Acquire) {
if let Some((_, handle)) = registry.remove_if(&undo_key, |_, h| h.id() == spawned_id) {
handle.abort();
}
return ArmOutcome::EpochClosed;
}
ArmOutcome::Armed
}
pub(crate) fn remove_watch(&self, parent_pid: u64, child_id: &WorkflowId) {
self.watches.remove(&(parent_pid, child_id.clone()));
}
pub(crate) fn remove_spawn_retry(&self, child_id: &WorkflowId) {
self.spawn_retries.remove(child_id);
}
pub(crate) fn remove_completion_retry(
&self,
lease: &CompletionRetryKey,
task: tokio::task::Id,
) {
self.completion_retries
.remove_if(lease, |_, handle| handle.id() == task);
}
pub(crate) fn abort_watch(&self, parent_pid: u64, child_id: &WorkflowId) {
if let Some((_, handle)) = self.watches.remove(&(parent_pid, child_id.clone())) {
handle.abort();
}
}
pub(crate) fn abort_watches_for_parent(&self, parent_pid: u64) {
self.watches.retain(|(pid, _), handle| {
if *pid == parent_pid {
handle.abort();
false
} else {
true
}
});
}
#[cfg(test)]
pub(crate) fn armed_watch_count(&self) -> usize {
self.watches.len()
}
#[cfg(test)]
pub(crate) fn armed_spawn_retry_count(&self) -> usize {
self.spawn_retries.len()
}
#[cfg(test)]
pub(crate) fn armed_completion_retry_count(&self) -> usize {
self.completion_retries.len()
}
pub(crate) fn shutdown(&self) {
self.gate_and_abort();
let runtime = {
let mut guard = match self.runtime.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
guard.take()
};
let Some(runtime) = runtime else {
return;
};
shutdown_runtime_and_join(runtime);
}
pub(crate) fn is_epoch_open(&self) -> bool {
!self.shutting_down.load(Ordering::Acquire)
}
fn gate_and_abort(&self) {
self.shutting_down.store(true, Ordering::Release);
self.watches.retain(|_, handle| {
handle.abort();
false
});
self.spawn_retries.retain(|_, handle| {
handle.abort();
false
});
self.completion_retries.retain(|_, handle| {
handle.abort();
false
});
}
pub(crate) fn begin_close(&self) {
self.gate_and_abort();
}
}
pub(crate) async fn sleep_backoff(current: &mut std::time::Duration, ceiling: std::time::Duration) {
tokio::time::sleep(*current).await;
let doubled = current.saturating_mul(2);
*current = if doubled > ceiling { ceiling } else { doubled };
}
fn shutdown_runtime_and_join(runtime: tokio::runtime::Runtime) {
match std::thread::Builder::new()
.name("aion-engine-tasks-shutdown".to_owned())
.spawn(move || drop(runtime))
{
Ok(joiner) => {
if joiner.join().is_err() {
tracing::error!("engine-task runtime shutdown thread panicked");
}
}
Err(error) => {
tracing::error!(
error = %error,
"could not spawn the engine-task runtime shutdown thread; the \
runtime was dropped on the calling thread instead"
);
}
}
}
impl Drop for EngineTaskRuntime {
fn drop(&mut self) {
self.gate_and_abort();
let runtime = {
let mut guard = match self.runtime.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
guard.take()
};
if let Some(runtime) = runtime {
runtime.shutdown_background();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use aion_core::{RunId, WorkflowId};
use super::{ArmOutcome, CompletionRetryKey, EngineTaskRuntime};
type TestResult = Result<(), Box<dyn std::error::Error>>;
struct DropFlag(Arc<AtomicBool>);
impl Drop for DropFlag {
fn drop(&mut self) {
self.0.store(true, Ordering::Release);
}
}
fn park_forever(flag: Arc<AtomicBool>) -> impl Future<Output = ()> + Send + 'static {
let guard = DropFlag(flag);
async move {
let _guard = guard;
loop {
tokio::time::sleep(Duration::from_secs(3600)).await;
}
}
}
#[test]
fn arming_is_idempotent_per_key() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let parent = 7;
let child = WorkflowId::new_v4();
let flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_watch(parent, child.clone(), park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
assert_eq!(
tasks.arm_watch(parent, child.clone(), park_forever(Arc::clone(&flag))),
ArmOutcome::AlreadyArmed
);
assert_eq!(tasks.armed_watch_count(), 1);
assert_eq!(
tasks.arm_watch(
parent,
WorkflowId::new_v4(),
park_forever(Arc::clone(&flag))
),
ArmOutcome::Armed
);
assert_eq!(tasks.armed_watch_count(), 2);
tasks.shutdown();
Ok(())
}
#[test]
fn abort_watch_disarms_a_single_key() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let child = WorkflowId::new_v4();
let other = WorkflowId::new_v4();
let flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_watch(3, child.clone(), park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
assert_eq!(
tasks.arm_watch(3, other.clone(), park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
tasks.abort_watch(3, &child);
assert_eq!(tasks.armed_watch_count(), 1);
assert_eq!(
tasks.arm_watch(3, other, park_forever(Arc::clone(&flag))),
ArmOutcome::AlreadyArmed
);
assert_eq!(
tasks.arm_watch(3, child, park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
tasks.shutdown();
Ok(())
}
#[test]
fn abort_for_parent_leaves_other_parents_armed() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let flag = Arc::new(AtomicBool::new(false));
for parent in [31, 31, 32] {
assert_eq!(
tasks.arm_watch(
parent,
WorkflowId::new_v4(),
park_forever(Arc::clone(&flag))
),
ArmOutcome::Armed
);
}
tasks.abort_watches_for_parent(31);
assert_eq!(tasks.armed_watch_count(), 1);
tasks.shutdown();
Ok(())
}
#[test]
fn a_release_from_a_foreign_task_cannot_evict_a_live_completion_retry() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let run = CompletionRetryKey {
workflow_id: WorkflowId::new_v4(),
run_id: RunId::new_v4(),
monitor_pid: 1,
};
let flag = Arc::new(AtomicBool::new(false));
let (foreign_tx, foreign_rx) = std::sync::mpsc::channel();
assert_eq!(
tasks.arm_spawn_retry(WorkflowId::new_v4(), async move {
let _ = foreign_tx.send(tokio::task::id());
}),
ArmOutcome::Armed
);
let foreign = foreign_rx.recv_timeout(Duration::from_secs(10))?;
let (live_tx, live_rx) = std::sync::mpsc::channel();
assert_eq!(
tasks.arm_completion_retry(run.clone(), {
let flag = Arc::clone(&flag);
async move {
let _ = live_tx.send(tokio::task::id());
park_forever(flag).await;
}
}),
ArmOutcome::Armed
);
let live = live_rx.recv_timeout(Duration::from_secs(10))?;
assert_ne!(
foreign, live,
"control: the two ids must differ, or the assertions below cannot tell the identity \
check from an unconditional removal"
);
assert_eq!(tasks.armed_completion_retry_count(), 1);
tasks.remove_completion_retry(&run, foreign);
assert_eq!(
tasks.armed_completion_retry_count(),
1,
"a release issued for a different task evicted the live retry's registration; the run \
now reads as unowned while a retry for it is still running, and the next arm for it \
would spawn a second writer of the same terminal"
);
tasks.remove_completion_retry(&run, live);
assert_eq!(
tasks.armed_completion_retry_count(),
0,
"the owning task's own release must clear its entry, or a finished retry would pin \
its run's key forever and no later exit for that run could ever arm"
);
tasks.shutdown();
Ok(())
}
#[test]
fn a_reopened_runs_successor_lease_can_arm_its_own_completion_retry() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let workflow_id = WorkflowId::new_v4();
let run_id = RunId::new_v4();
let predecessor = CompletionRetryKey {
workflow_id: workflow_id.clone(),
run_id: run_id.clone(),
monitor_pid: 11,
};
let successor = CompletionRetryKey {
workflow_id,
run_id,
monitor_pid: 12,
};
assert_ne!(
predecessor.monitor_pid, successor.monitor_pid,
"control: the fixture must have built two DIFFERENT leases — same workflow, same \
run, different monitor pid — or there is no reopen here to test"
);
assert_eq!(
(&predecessor.workflow_id, &predecessor.run_id),
(&successor.workflow_id, &successor.run_id),
"control: and they must be the SAME run, or this is two unrelated workflows and the \
collision the reopen causes never arises"
);
let predecessor_flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_completion_retry(
predecessor.clone(),
park_forever(Arc::clone(&predecessor_flag))
),
ArmOutcome::Armed,
"the superseded lease's retry is armed and still sleeping on its backoff"
);
let duplicate_flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_completion_retry(predecessor, park_forever(Arc::clone(&duplicate_flag))),
ArmOutcome::AlreadyArmed,
"control: one LEASE must still never hold two retries — if this is `Armed` the pid \
did not narrow the key, it replaced it, and one run's terminal has two writers"
);
let successor_flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_completion_retry(successor, park_forever(Arc::clone(&successor_flag))),
ArmOutcome::Armed,
"THE DECISIVE ONE: a reopened run's successor lease was refused as `AlreadyArmed` by \
its own superseded predecessor, whose retry then stood down without writing — so \
the successor's terminal was never recorded and the run projected Running forever"
);
assert_eq!(
tasks.armed_completion_retry_count(),
2,
"both leases hold a registration: the superseded one until it stands down, the \
successor until it writes"
);
tasks.shutdown();
Ok(())
}
#[test]
fn shutdown_gates_new_arms_and_awaits_aborted_tasks() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let watch_flag = Arc::new(AtomicBool::new(false));
let retry_flag = Arc::new(AtomicBool::new(false));
let completion_flag = Arc::new(AtomicBool::new(false));
let child = WorkflowId::new_v4();
let run = CompletionRetryKey {
workflow_id: WorkflowId::new_v4(),
run_id: RunId::new_v4(),
monitor_pid: 1,
};
assert_eq!(
tasks.arm_watch(9, child.clone(), park_forever(Arc::clone(&watch_flag))),
ArmOutcome::Armed
);
assert_eq!(
tasks.arm_spawn_retry(child.clone(), park_forever(Arc::clone(&retry_flag))),
ArmOutcome::Armed
);
assert_eq!(
tasks.arm_completion_retry(run.clone(), park_forever(Arc::clone(&completion_flag))),
ArmOutcome::Armed
);
assert_eq!(tasks.armed_spawn_retry_count(), 1);
assert_eq!(tasks.armed_completion_retry_count(), 1);
tasks.shutdown();
assert!(
watch_flag.load(Ordering::Acquire),
"watcher task must be fully dropped before the epoch closes"
);
assert!(
retry_flag.load(Ordering::Acquire),
"spawn-retry task must be fully dropped before the epoch closes"
);
assert!(
completion_flag.load(Ordering::Acquire),
"completion-retry task must be fully dropped before the epoch closes"
);
assert_eq!(tasks.armed_watch_count(), 0);
assert_eq!(tasks.armed_spawn_retry_count(), 0);
assert_eq!(tasks.armed_completion_retry_count(), 0);
assert_eq!(
tasks.arm_watch(9, child.clone(), park_forever(Arc::clone(&watch_flag))),
ArmOutcome::EpochClosed
);
assert_eq!(
tasks.arm_spawn_retry(child, park_forever(Arc::clone(&retry_flag))),
ArmOutcome::EpochClosed
);
assert_eq!(
tasks.arm_completion_retry(run, park_forever(Arc::clone(&completion_flag))),
ArmOutcome::EpochClosed
);
Ok(())
}
#[test]
fn dropping_the_registry_aborts_completion_retries_too() -> TestResult {
let flag = Arc::new(AtomicBool::new(false));
{
let tasks = EngineTaskRuntime::new()?;
assert_eq!(
tasks.arm_completion_retry(
CompletionRetryKey {
workflow_id: WorkflowId::new_v4(),
run_id: RunId::new_v4(),
monitor_pid: 1,
},
park_forever(Arc::clone(&flag)),
),
ArmOutcome::Armed
);
assert_eq!(tasks.armed_completion_retry_count(), 1);
}
let aborted = (0..500).any(|_| {
if flag.load(Ordering::Acquire) {
return true;
}
std::thread::sleep(std::time::Duration::from_millis(10));
false
});
assert!(
aborted,
"dropping the registry must abort the armed completion retry"
);
Ok(())
}
#[tokio::test]
async fn shutdown_is_safe_from_inside_a_host_async_context() -> TestResult {
let tasks = EngineTaskRuntime::new()?;
let flag = Arc::new(AtomicBool::new(false));
assert_eq!(
tasks.arm_watch(11, WorkflowId::new_v4(), park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
tasks.shutdown();
assert!(flag.load(Ordering::Acquire));
Ok(())
}
#[tokio::test]
async fn drop_backstop_aborts_without_blocking() -> TestResult {
let flag = Arc::new(AtomicBool::new(false));
{
let tasks = EngineTaskRuntime::new()?;
assert_eq!(
tasks.arm_watch(12, WorkflowId::new_v4(), park_forever(Arc::clone(&flag))),
ArmOutcome::Armed
);
}
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while !flag.load(Ordering::Acquire) {
if std::time::Instant::now() > deadline {
return Err("drop backstop never aborted the armed task".into());
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
Ok(())
}
}