use chrono::{DateTime, Utc};
use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{Duration, Instant};
use tokio::sync::{Notify, watch};
use tracing::{error, info, warn};
use aion_core::{ActivityId, Payload, WorkflowId};
use aion_proto::{ProtoHeartbeat, WireError};
use crate::error::ServerError;
use crate::shutdown::DrainState;
use crate::worker::dispatch::{
ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink,
};
use crate::worker::envelope::CompletionToken;
use crate::worker::registry::{ConnectedWorkerRegistry, WorkerId};
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct InFlightActivity {
pub workflow_id: WorkflowId,
pub activity_id: ActivityId,
pub attempt: u32,
pub completion_token: CompletionToken,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TaskLiveness {
pub worker_id: WorkerId,
pub workflow_id: WorkflowId,
pub activity_id: ActivityId,
pub attempt: u32,
pub completion_token: CompletionToken,
pub heartbeat_window: Duration,
pub last_heartbeat_at: Instant,
pub last_progress: Option<Payload>,
pub last_progress_at: Option<DateTime<Utc>>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct HeartbeatUpdate {
pub liveness: TaskLiveness,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct LostWorkerReport {
pub worker_id: WorkerId,
pub tasks: Vec<InFlightActivity>,
pub task_queue: Option<String>,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
struct TaskKey(WorkerId, WorkflowId, ActivityId);
#[derive(Debug, Default)]
struct HeartbeatState {
tasks: HashMap<TaskKey, TaskLiveness>,
connections: HashMap<WorkerId, Instant>,
reachability: HashMap<WorkerId, Reachability>,
}
pub(crate) const DISPATCH_PROBATION_PINGS: u32 = 2;
#[derive(Clone, Copy, Debug, Default)]
struct Reachability {
consecutive_answers: u32,
proved_at: Option<Instant>,
ever_proved: bool,
}
impl Reachability {
fn is_proved(self, now: Instant, window: Duration) -> bool {
self.consecutive_answers >= DISPATCH_PROBATION_PINGS
&& self.proved_at.is_some_and(|proved_at| {
now.checked_duration_since(proved_at)
.is_none_or(|elapsed| elapsed <= window)
})
}
fn exclusion(self) -> DispatchExclusion {
if self.ever_proved {
DispatchExclusion::ReachabilityLost
} else {
DispatchExclusion::OpeningProbation {
answers: self.consecutive_answers,
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DispatchExclusion {
OpeningProbation {
answers: u32,
},
ReachabilityLost,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ExcludedWorker {
pub worker_id: WorkerId,
pub exclusion: DispatchExclusion,
}
#[derive(Clone, Debug)]
pub struct HeartbeatTracker {
heartbeat_window: Duration,
inner: Arc<Mutex<HeartbeatState>>,
empty: Arc<Notify>,
notes_held_since: DateTime<Utc>,
}
impl HeartbeatTracker {
#[must_use]
pub fn new(heartbeat_window: Duration) -> Self {
Self {
heartbeat_window,
inner: Arc::new(Mutex::new(HeartbeatState::default())),
empty: Arc::new(Notify::new()),
notes_held_since: Utc::now(),
}
}
#[must_use]
pub const fn notes_held_since(&self) -> DateTime<Utc> {
self.notes_held_since
}
pub fn register_connection(
&self,
worker_id: WorkerId,
now: Instant,
) -> Result<(), ServerError> {
let mut state = self.state()?;
state.connections.insert(worker_id, now);
state
.reachability
.insert(worker_id, Reachability::default());
Ok(())
}
pub fn record_connection_activity(
&self,
worker_id: WorkerId,
now: Instant,
) -> Result<bool, ServerError> {
let mut state = self.state()?;
let Some(last_activity) = state.connections.get_mut(&worker_id) else {
return Ok(false);
};
*last_activity = now;
Ok(true)
}
pub fn record_dispatch_reachability(
&self,
worker_id: WorkerId,
now: Instant,
) -> Result<bool, ServerError> {
let mut state = self.state()?;
let Some(last_activity) = state.connections.get_mut(&worker_id) else {
return Ok(false);
};
*last_activity = now;
let standing = state.reachability.entry(worker_id).or_default();
standing.consecutive_answers = standing.consecutive_answers.saturating_add(1);
standing.proved_at = Some(now);
if standing.consecutive_answers >= DISPATCH_PROBATION_PINGS {
standing.ever_proved = true;
}
Ok(true)
}
pub fn record_dispatch_unreachable(&self, worker_id: WorkerId) -> Result<bool, ServerError> {
let mut state = self.state()?;
if !state.connections.contains_key(&worker_id) {
return Ok(false);
}
let standing = state.reachability.entry(worker_id).or_default();
standing.consecutive_answers = 0;
standing.proved_at = None;
Ok(true)
}
pub fn is_dispatch_reachable(
&self,
worker_id: WorkerId,
now: Instant,
) -> Result<bool, ServerError> {
let state = self.state()?;
Ok(state
.reachability
.get(&worker_id)
.is_some_and(|standing| standing.is_proved(now, self.heartbeat_window)))
}
pub fn unreachable_workers(&self, now: Instant) -> Result<Vec<ExcludedWorker>, ServerError> {
let state = self.state()?;
let mut workers = state
.reachability
.iter()
.filter(|(_, standing)| !standing.is_proved(now, self.heartbeat_window))
.map(|(worker_id, standing)| ExcludedWorker {
worker_id: *worker_id,
exclusion: standing.exclusion(),
})
.collect::<Vec<_>>();
workers.sort_unstable_by_key(|excluded| excluded.worker_id);
Ok(workers)
}
pub fn unregister_connection(&self, worker_id: WorkerId) -> Result<(), ServerError> {
let mut state = self.state()?;
state.connections.remove(&worker_id);
state.reachability.remove(&worker_id);
Ok(())
}
pub fn track_task(
&self,
worker_id: WorkerId,
task: InFlightActivity,
now: Instant,
) -> Result<(), ServerError> {
let key = TaskKey::new(
worker_id,
task.workflow_id.clone(),
task.activity_id.clone(),
);
let liveness = TaskLiveness {
worker_id,
workflow_id: task.workflow_id,
activity_id: task.activity_id,
attempt: task.attempt,
completion_token: task.completion_token,
heartbeat_window: self.heartbeat_window,
last_heartbeat_at: now,
last_progress: None,
last_progress_at: None,
};
let mut state = self.state()?;
state.tasks.insert(key, liveness);
state.connections.insert(worker_id, now);
Ok(())
}
pub fn complete_task(
&self,
worker_id: WorkerId,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
) -> Result<bool, ServerError> {
let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
let (was_tracked, became_empty) = {
let mut state = self.state()?;
let was_tracked = state.tasks.remove(&key).is_some();
(was_tracked, state.tasks.is_empty())
};
if became_empty {
self.empty.notify_waiters();
}
Ok(was_tracked)
}
pub fn is_tracked(
&self,
worker_id: WorkerId,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
) -> Result<bool, ServerError> {
let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
Ok(self.state()?.tasks.contains_key(&key))
}
pub fn record_liveness(
&self,
worker_id: WorkerId,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
now: Instant,
) -> Result<bool, ServerError> {
let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
let mut state = self.state()?;
if !state.tasks.contains_key(&key) {
return Ok(false);
}
if let Some(last_activity) = state.connections.get_mut(&worker_id) {
*last_activity = now;
}
let Some(liveness) = state.tasks.get_mut(&key) else {
return Ok(false);
};
liveness.last_heartbeat_at = now;
Ok(true)
}
#[must_use]
pub const fn heartbeat_window(&self) -> Duration {
self.heartbeat_window
}
pub fn in_flight_count(&self) -> Result<usize, ServerError> {
Ok(self.state()?.tasks.len())
}
pub fn record_heartbeat(
&self,
worker_id: WorkerId,
heartbeat: ProtoHeartbeat,
now: Instant,
) -> Result<HeartbeatUpdate, ServerError> {
let decoded = DecodedHeartbeat::try_from(heartbeat)?;
let key = TaskKey::new(worker_id, decoded.workflow_id, decoded.activity_id);
let mut state = self.state()?;
if !state.tasks.contains_key(&key) {
return Err(wire_error("heartbeat task is not in flight"));
}
if let Some(last_activity) = state.connections.get_mut(&worker_id) {
*last_activity = now;
}
let Some(liveness) = state.tasks.get_mut(&key) else {
return Err(wire_error("heartbeat task is not in flight"));
};
liveness.last_heartbeat_at = now;
if decoded.progress.is_some() {
liveness.last_progress = decoded.progress;
liveness.last_progress_at = Some(Utc::now());
}
Ok(HeartbeatUpdate {
liveness: liveness.clone(),
})
}
pub fn is_live(
&self,
worker_id: WorkerId,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
now: Instant,
) -> Result<bool, ServerError> {
let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
let state = self.state()?;
let Some(liveness) = state.tasks.get(&key) else {
return Err(wire_error("heartbeat task is not in flight"));
};
Ok(!is_expired(liveness, now))
}
pub fn expired_workers(&self, now: Instant) -> Result<Vec<WorkerId>, ServerError> {
let state = self.state()?;
let mut seen = HashSet::new();
let mut workers = Vec::new();
for (worker_id, last_activity) in &state.connections {
if now
.checked_duration_since(*last_activity)
.is_some_and(|elapsed| elapsed > self.heartbeat_window)
&& seen.insert(*worker_id)
{
workers.push(*worker_id);
}
}
for liveness in state.tasks.values() {
if is_expired(liveness, now) && seen.insert(liveness.worker_id) {
workers.push(liveness.worker_id);
}
}
workers.sort_unstable();
Ok(workers)
}
pub fn fail_expired_workers(
&self,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
now: Instant,
) -> Result<Vec<LostWorkerReport>, ServerError> {
let mut reports = Vec::new();
for worker_id in self.expired_workers(now)? {
let report = self.fail_lost_worker(worker_id, registry, sink)?;
reports.push(report);
}
Ok(reports)
}
pub fn fail_disconnected_worker(
&self,
worker_id: WorkerId,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
) -> Result<LostWorkerReport, ServerError> {
self.fail_lost_worker(worker_id, registry, sink)
}
pub fn fail_all_in_flight_workers(
&self,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
) -> Result<Vec<LostWorkerReport>, ServerError> {
let worker_ids = {
let state = self.state()?;
let mut worker_ids = state
.tasks
.values()
.map(|liveness| liveness.worker_id)
.collect::<HashSet<_>>()
.into_iter()
.collect::<Vec<_>>();
worker_ids.sort_unstable();
worker_ids
};
let mut reports = Vec::new();
for worker_id in worker_ids {
let report = self.fail_lost_worker(worker_id, registry, sink)?;
if !report.tasks.is_empty() {
reports.push(report);
}
}
self.empty.notify_waiters();
Ok(reports)
}
pub fn park_disconnected_worker(
&self,
worker_id: WorkerId,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
) -> Result<LostWorkerReport, ServerError> {
self.park_lost_worker(
worker_id,
registry,
sink,
aion_core::WorkerDeathReason::Disconnect,
)
}
pub fn park_all_in_flight_workers(
&self,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
) -> Result<Vec<LostWorkerReport>, ServerError> {
let worker_ids = {
let state = self.state()?;
let mut worker_ids = state
.tasks
.values()
.map(|liveness| liveness.worker_id)
.collect::<HashSet<_>>()
.into_iter()
.collect::<Vec<_>>();
worker_ids.sort_unstable();
worker_ids
};
let mut reports = Vec::new();
for worker_id in worker_ids {
let report = self.park_lost_worker(
worker_id,
registry,
sink,
aion_core::WorkerDeathReason::Timeout,
)?;
if !report.tasks.is_empty() {
reports.push(report);
}
}
self.empty.notify_waiters();
Ok(reports)
}
fn park_lost_worker(
&self,
worker_id: WorkerId,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
reason: aion_core::WorkerDeathReason,
) -> Result<LostWorkerReport, ServerError> {
let task_queue = task_queue_of(registry, worker_id);
registry.deregister_with_reason(worker_id, reason)?;
self.state()?.connections.remove(&worker_id);
let tasks = self.remove_worker_tasks(worker_id)?;
for task in &tasks {
sink.park_activity(&task.workflow_id, &task.activity_id)?;
info!(
worker_id = ?worker_id,
workflow_id = %task.workflow_id,
activity_id = %task.activity_id,
"activity parked for restart recovery"
);
}
Ok(LostWorkerReport {
worker_id,
tasks,
task_queue,
})
}
fn fail_lost_worker(
&self,
worker_id: WorkerId,
registry: &ConnectedWorkerRegistry,
sink: &impl ActivityCompletionSink,
) -> Result<LostWorkerReport, ServerError> {
let task_queue = task_queue_of(registry, worker_id);
registry.deregister_with_reason(worker_id, aion_core::WorkerDeathReason::Timeout)?;
self.state()?.connections.remove(&worker_id);
let tasks = self.remove_worker_tasks(worker_id)?;
for task in &tasks {
sink.complete_activity(ActivityCompletion {
workflow_id: task.workflow_id.clone(),
activity_id: task.activity_id.clone(),
run_id: None,
completion_token: task.completion_token.clone(),
outcome: ActivityCompletionOutcome::WorkerLost { worker_id },
})?;
}
Ok(LostWorkerReport {
worker_id,
tasks,
task_queue,
})
}
fn remove_worker_tasks(
&self,
worker_id: WorkerId,
) -> Result<Vec<InFlightActivity>, ServerError> {
let mut state = self.state()?;
let keys = state
.tasks
.keys()
.filter(|key| key.worker_id() == worker_id)
.cloned()
.collect::<Vec<_>>();
let mut tasks = Vec::with_capacity(keys.len());
for key in keys {
if let Some(liveness) = state.tasks.remove(&key) {
tasks.push(InFlightActivity {
workflow_id: liveness.workflow_id,
activity_id: liveness.activity_id,
attempt: liveness.attempt,
completion_token: liveness.completion_token,
});
}
}
Ok(tasks)
}
pub(in crate::worker) fn attempt_entries(
&self,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
attempt: u32,
) -> Result<Vec<TaskLiveness>, ServerError> {
Ok(self
.state()?
.tasks
.values()
.filter(|liveness| {
&liveness.workflow_id == workflow_id
&& &liveness.activity_id == activity_id
&& liveness.attempt == attempt
})
.cloned()
.collect())
}
pub fn in_flight_for_workflow(
&self,
workflow_id: &WorkflowId,
) -> Result<Vec<TaskLiveness>, ServerError> {
Ok(self
.state()?
.tasks
.values()
.filter(|liveness| &liveness.workflow_id == workflow_id)
.cloned()
.collect())
}
fn state(&self) -> Result<MutexGuard<'_, HeartbeatState>, ServerError> {
self.inner
.lock()
.map_err(|_| ServerError::lock_poisoned("worker heartbeat tracker"))
}
#[cfg(test)]
pub(crate) fn poison_for_tests(&self) {
let inner = Arc::clone(&self.inner);
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(move || {
let _guard = inner.lock();
std::panic::resume_unwind(Box::new(
"poisoning the heartbeat tracker for a fail-open test",
));
}));
}
}
#[must_use]
pub fn sweep_interval(heartbeat_window: Duration) -> Duration {
const MINIMUM_PERIOD: Duration = Duration::from_millis(1);
const TARGET_FLOOR: Duration = Duration::from_secs(1);
let ceiling = heartbeat_window.max(MINIMUM_PERIOD);
(heartbeat_window / 4).clamp(TARGET_FLOOR.min(ceiling), ceiling)
}
pub struct HeartbeatSweeper<S> {
tracker: HeartbeatTracker,
registry: ConnectedWorkerRegistry,
sink: S,
drain: DrainState,
heartbeat_window: Duration,
interval: Duration,
queue_state: crate::worker::QueueServiceState,
}
impl<S> HeartbeatSweeper<S>
where
S: ActivityCompletionSink + Send + Sync + 'static,
{
#[must_use]
pub fn new(
tracker: HeartbeatTracker,
registry: ConnectedWorkerRegistry,
sink: S,
drain: DrainState,
heartbeat_window: Duration,
) -> Self {
let interval = sweep_interval(heartbeat_window);
Self {
tracker,
registry,
sink,
drain,
heartbeat_window,
interval,
queue_state: crate::worker::QueueServiceState::default(),
}
}
#[must_use]
pub fn with_queue_state(mut self, queue_state: crate::worker::QueueServiceState) -> Self {
self.queue_state = queue_state;
self
}
pub async fn run(self, mut shutdown: watch::Receiver<bool>) {
info!(
sweep_interval_ms = self.interval.as_millis(),
heartbeat_window_ms = self.heartbeat_window.as_millis(),
"worker heartbeat sweeper started"
);
let mut ticks = tokio::time::interval(self.interval);
ticks.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = ticks.tick() => {
if *shutdown.borrow() {
break;
}
self.sweep_once(Instant::now());
}
changed = shutdown.changed() => {
if changed.is_err() || *shutdown.borrow() {
break;
}
}
}
}
info!("worker heartbeat sweeper stopped");
}
fn sweep_once(&self, now: Instant) {
let reports = match self
.tracker
.fail_expired_workers(&self.registry, &self.sink, now)
{
Ok(reports) => reports,
Err(sweep_error) => {
error!(
error = %sweep_error,
"heartbeat expiry sweep failed; retrying next tick"
);
return;
}
};
for report in &reports {
let task_queue = report.task_queue.as_deref().unwrap_or("<unregistered>");
let parked = report
.task_queue
.as_deref()
.map(|queue| self.queue_state.parked_on_queue(queue))
.transpose()
.unwrap_or_else(|error| {
error!(%error, "could not read parked-dispatch count for a reaped worker");
None
})
.unwrap_or(0);
if report.tasks.is_empty() {
warn!(
worker_id = ?report.worker_id,
task_queue,
parked_dispatches = parked,
heartbeat_window_ms = self.heartbeat_window.as_millis(),
"idle worker connection lease expired; worker deregistered"
);
} else {
warn!(
worker_id = ?report.worker_id,
task_queue,
parked_dispatches = parked,
failed_tasks = report.tasks.len(),
heartbeat_window_ms = self.heartbeat_window.as_millis(),
"worker heartbeat window expired with in-flight activities; \
deregistered and surfaced as transport losses, to be \
re-dispatched attempt-neutrally"
);
}
}
if !reports.is_empty() {
self.drain.notify_activity_drained();
}
}
}
impl TaskKey {
fn new(worker_id: WorkerId, workflow_id: WorkflowId, activity_id: ActivityId) -> Self {
Self(worker_id, workflow_id, activity_id)
}
const fn worker_id(&self) -> WorkerId {
self.0
}
}
struct DecodedHeartbeat {
workflow_id: WorkflowId,
activity_id: ActivityId,
progress: Option<Payload>,
}
impl TryFrom<ProtoHeartbeat> for DecodedHeartbeat {
type Error = ServerError;
fn try_from(value: ProtoHeartbeat) -> Result<Self, Self::Error> {
let workflow_id = value
.workflow_id
.ok_or_else(|| wire_error("heartbeat workflow id is missing"))
.and_then(|id| WorkflowId::try_from(id).map_err(ServerError::from))?;
let activity_id = value
.activity_id
.ok_or_else(|| wire_error("heartbeat activity id is missing"))
.map(ActivityId::from)?;
let progress = value
.progress
.map(Payload::try_from)
.transpose()
.map_err(ServerError::from)?;
Ok(Self {
workflow_id,
activity_id,
progress,
})
}
}
fn task_queue_of(registry: &ConnectedWorkerRegistry, worker_id: WorkerId) -> Option<String> {
registry
.worker_by_id(worker_id)
.ok()
.flatten()
.map(|handle| handle.task_queue().to_owned())
}
fn is_expired(liveness: &TaskLiveness, now: Instant) -> bool {
now.checked_duration_since(liveness.last_heartbeat_at)
.is_some_and(|elapsed| elapsed > liveness.heartbeat_window)
}
fn wire_error(message: &'static str) -> ServerError {
ServerError::Wire {
wire: WireError::backend(message),
}
}
#[cfg(test)]
mod reachability_tests {
use std::time::{Duration, Instant};
use super::{
DISPATCH_PROBATION_PINGS, DispatchExclusion, ExcludedWorker, HeartbeatTracker, ServerError,
WorkerId,
};
const WINDOW: Duration = Duration::from_secs(30);
type TestResult = Result<(), ServerError>;
fn tracker_with_worker(now: Instant) -> Result<(HeartbeatTracker, WorkerId), ServerError> {
let tracker = HeartbeatTracker::new(WINDOW);
let worker = WorkerId::from_value(1);
tracker.register_connection(worker, now)?;
Ok((tracker, worker))
}
fn serve_probation(
tracker: &HeartbeatTracker,
worker: WorkerId,
at: Instant,
) -> Result<(), ServerError> {
for _ in 0..DISPATCH_PROBATION_PINGS {
assert!(
tracker.record_dispatch_reachability(worker, at)?,
"the worker must still be tracked while it serves its probation"
);
}
Ok(())
}
#[test]
fn an_inbound_frame_cannot_prove_dispatch_reachability() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
serve_probation(&tracker, worker, start)?;
assert!(
tracker.is_dispatch_reachable(worker, start)?,
"precondition: the worker is eligible before the connection goes one-way"
);
let much_later = start + WINDOW * 4;
assert!(
tracker.record_connection_activity(worker, much_later)?,
"the worker is still tracked"
);
assert!(
!tracker.is_dispatch_reachable(worker, much_later)?,
"a pump beat must NOT make a worker the server cannot push to look reachable"
);
assert_eq!(
tracker.unreachable_workers(much_later)?,
vec![ExcludedWorker {
worker_id: worker,
exclusion: DispatchExclusion::ReachabilityLost,
}],
"the worker must be named unreachable however alive its process looks"
);
Ok(())
}
#[test]
fn an_answered_ping_does_prove_dispatch_reachability() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
let much_later = start + WINDOW * 4;
serve_probation(&tracker, worker, much_later)?;
assert!(
tracker.is_dispatch_reachable(worker, much_later)?,
"answered pings are the one thing that proves the push leg works"
);
assert!(
tracker.unreachable_workers(much_later)?.is_empty(),
"a worker answering pings is never unreachable"
);
Ok(())
}
#[test]
fn registration_opens_a_probation_and_does_not_grant_eligibility() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
assert!(
!tracker.is_dispatch_reachable(worker, start)?,
"a brand-new connection has proved nothing about the push leg"
);
assert_eq!(
tracker.unreachable_workers(start)?,
vec![ExcludedWorker {
worker_id: worker,
exclusion: DispatchExclusion::OpeningProbation { answers: 0 },
}],
"a worker serving its probation is carried in the census as unreachable"
);
Ok(())
}
#[test]
fn a_part_served_probation_is_a_probation_and_not_a_reachability_failure() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
const {
assert!(
DISPATCH_PROBATION_PINGS > 1,
"this test is only meaningful while the probation takes more than one answer"
);
}
assert!(tracker.record_dispatch_reachability(worker, start)?);
assert_eq!(
tracker.unreachable_workers(start)?,
vec![ExcludedWorker {
worker_id: worker,
exclusion: DispatchExclusion::OpeningProbation { answers: 1 },
}],
"a worker that has answered part of its opening probation is still excluded, but it \
must not be described as one the server cannot reach — it answered"
);
Ok(())
}
#[test]
fn losing_held_eligibility_is_reported_as_a_loss_not_as_a_fresh_probation() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
serve_probation(&tracker, worker, start)?;
assert!(
tracker.is_dispatch_reachable(worker, start)?,
"precondition: eligibility was genuinely held before it was lost"
);
assert!(
tracker.record_dispatch_unreachable(worker)?,
"the worker is still tracked when its ping fails"
);
assert_eq!(
tracker.unreachable_workers(start)?,
vec![ExcludedWorker {
worker_id: worker,
exclusion: DispatchExclusion::ReachabilityLost,
}],
"a failed ping on a worker that HAD eligibility is an incident, and must not be \
filed as the ordinary probation every fresh connection serves"
);
Ok(())
}
#[test]
fn a_reconnect_starts_a_fresh_probation_not_a_lost_eligibility() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
serve_probation(&tracker, worker, start)?;
tracker.unregister_connection(worker)?;
tracker.register_connection(worker, start)?;
assert_eq!(
tracker.unreachable_workers(start)?,
vec![ExcludedWorker {
worker_id: worker,
exclusion: DispatchExclusion::OpeningProbation { answers: 0 },
}],
"nothing earned on the old connection carries across to the new one"
);
Ok(())
}
#[test]
fn one_ping_short_of_the_probation_earns_nothing() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
for _ in 0..DISPATCH_PROBATION_PINGS - 1 {
assert!(tracker.record_dispatch_reachability(worker, start)?);
assert!(
!tracker.is_dispatch_reachable(worker, start)?,
"eligibility must not be granted before the probation is served in full"
);
}
assert!(tracker.record_dispatch_reachability(worker, start)?);
assert!(
tracker.is_dispatch_reachable(worker, start)?,
"the ping that completes the probation must grant eligibility — otherwise this test \
would pass on a tracker that never grants it at all"
);
Ok(())
}
#[test]
fn a_failed_probe_withdraws_eligibility_at_once_and_restarts_the_probation() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
serve_probation(&tracker, worker, start)?;
assert!(
tracker.is_dispatch_reachable(worker, start)?,
"precondition"
);
assert!(
tracker.record_dispatch_unreachable(worker)?,
"the worker is still tracked"
);
assert!(
!tracker.is_dispatch_reachable(worker, start)?,
"a failed probe withdraws eligibility on the spot, inside the window"
);
assert!(tracker.record_dispatch_reachability(worker, start)?);
assert!(
!tracker.is_dispatch_reachable(worker, start)?,
"a single answer after a failure must not restore eligibility"
);
Ok(())
}
#[test]
fn a_link_that_answers_every_other_probe_never_becomes_eligible() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
for probe in 0..DISPATCH_PROBATION_PINGS * 10 {
let now = start + Duration::from_millis(u64::from(probe));
if probe % 2 == 0 {
assert!(tracker.record_dispatch_reachability(worker, now)?);
} else {
assert!(tracker.record_dispatch_unreachable(worker)?);
}
assert!(
!tracker.is_dispatch_reachable(worker, now)?,
"a flapping link must never hold dispatch eligibility, at any probe (probe {probe})"
);
}
let now = start + Duration::from_secs(1);
serve_probation(&tracker, worker, now)?;
assert!(
tracker.is_dispatch_reachable(worker, now)?,
"consecutive answers must still earn eligibility"
);
Ok(())
}
#[test]
fn reachability_goes_stale_once_the_window_passes() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
serve_probation(&tracker, worker, start)?;
assert!(
tracker.is_dispatch_reachable(worker, start + WINDOW)?,
"still inside the window"
);
assert!(
!tracker.is_dispatch_reachable(worker, start + WINDOW + Duration::from_millis(1))?,
"one millisecond past the window is stale"
);
Ok(())
}
#[test]
fn an_unregistered_worker_is_never_reachable_and_cannot_be_resurrected() -> TestResult {
let start = Instant::now();
let (tracker, worker) = tracker_with_worker(start)?;
tracker.unregister_connection(worker)?;
assert!(
!tracker.is_dispatch_reachable(worker, start)?,
"a deregistered worker is not reachable"
);
assert!(
!tracker.record_dispatch_reachability(worker, start)?,
"a late pong must not resurrect a deregistered worker"
);
assert!(
!tracker.record_dispatch_unreachable(worker)?,
"a late probe FAILURE must not resurrect a deregistered worker either — the reset \
path allocates an entry, so it has to refuse an untracked worker as firmly as the \
success path does"
);
assert!(
tracker.unreachable_workers(start)?.is_empty(),
"an untracked worker is not carried in the census either"
);
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use aion_core::ContentType;
use aion_proto::{ProtoActivityId, ProtoPayload, ProtoWorkflowId};
use serde_json::json;
use uuid::Uuid;
use crate::worker::registry::WorkerRegistration;
use super::*;
#[derive(Default)]
struct RecordingSink {
completions: Mutex<Vec<ActivityCompletion>>,
parks: Mutex<Vec<(WorkflowId, ActivityId)>>,
}
impl ActivityCompletionSink for RecordingSink {
fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
self.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.push(completion);
Ok(())
}
fn park_activity(
&self,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
) -> Result<(), ServerError> {
self.parks
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.push((workflow_id.clone(), activity_id.clone()));
Ok(())
}
}
fn workflow_id() -> WorkflowId {
WorkflowId::new(Uuid::nil())
}
fn activity_id(position: u64) -> ActivityId {
ActivityId::from_sequence_position(position)
}
fn payload(value: &serde_json::Value) -> Result<Payload, Box<dyn std::error::Error>> {
Ok(Payload::from_json(value)?)
}
fn heartbeat(
workflow_id: WorkflowId,
activity_id: ActivityId,
progress: Option<Payload>,
) -> ProtoHeartbeat {
ProtoHeartbeat {
workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
activity_id: Some(ProtoActivityId::from(activity_id)),
progress: progress.map(ProtoPayload::from),
}
}
fn registry_with_worker()
-> Result<(ConnectedWorkerRegistry, WorkerRegistration, WorkerId), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = tokio::sync::mpsc::channel(1);
let activity_types = [String::from("charge-card")];
let registration = registry.register("tenant-a", activity_types.iter(), tx)?;
let worker_id = registration
.worker_id()
.ok_or_else(|| ServerError::lock_poisoned("test worker registration"))?;
Ok((registry, registration, worker_id))
}
#[test]
fn heartbeat_refresh_keeps_task_live_across_window() -> Result<(), Box<dyn std::error::Error>> {
let window = Duration::from_secs(5);
let tracker = HeartbeatTracker::new(window);
let worker_id = WorkerIdForTest::registered()?;
let workflow_id = workflow_id();
let activity_id = activity_id(10);
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
assert!(tracker.is_live(worker_id, &workflow_id, &activity_id, start + window)?);
let progress = payload(&json!({"percent": 50}))?;
let update = tracker.record_heartbeat(
worker_id,
heartbeat(
workflow_id.clone(),
activity_id.clone(),
Some(progress.clone()),
),
start + window,
)?;
assert_eq!(update.liveness.last_progress, Some(progress));
assert!(tracker.is_live(
worker_id,
&workflow_id,
&activity_id,
start + window + window
)?);
assert!(tracker.expired_workers(start + window + window)?.is_empty());
Ok(())
}
#[test]
fn missed_heartbeat_deregisters_worker_and_fails_in_flight_once()
-> Result<(), Box<dyn std::error::Error>> {
let (registry, _registration, worker_id) = registry_with_worker()?;
let sink = RecordingSink::default();
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let workflow_id = workflow_id();
let activity_id = activity_id(11);
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
let reports =
tracker.fail_expired_workers(®istry, &sink, start + Duration::from_secs(6))?;
assert_eq!(reports.len(), 1);
assert_eq!(reports[0].worker_id, worker_id);
assert_eq!(reports[0].tasks.len(), 1);
assert!(
registry
.workers_for("tenant-a", "default", "charge-card", None)?
.is_empty()
);
let second = tracker.fail_disconnected_worker(worker_id, ®istry, &sink)?;
assert!(second.tasks.is_empty());
let completions = sink
.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
assert_eq!(completions.len(), 1);
assert_eq!(completions[0].workflow_id, workflow_id);
assert_eq!(completions[0].activity_id, activity_id);
match &completions[0].outcome {
ActivityCompletionOutcome::WorkerLost { worker_id: lost } => {
assert_eq!(*lost, worker_id);
}
other => {
return Err(format!("expected a lost-worker outcome, got {other:?}").into());
}
}
Ok(())
}
#[test]
fn disconnected_worker_fails_each_in_flight_task_once() -> Result<(), Box<dyn std::error::Error>>
{
let (registry, _registration, worker_id) = registry_with_worker()?;
let sink = RecordingSink::default();
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let workflow_id = workflow_id();
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id(21),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id,
activity_id: activity_id(22),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
let report = tracker.fail_disconnected_worker(worker_id, ®istry, &sink)?;
assert_eq!(report.tasks.len(), 2);
assert!(
registry
.workers_for("tenant-a", "default", "charge-card", None)?
.is_empty()
);
let completions = sink
.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
assert_eq!(completions.len(), 2);
assert!(completions.iter().all(|completion| matches!(
&completion.outcome,
ActivityCompletionOutcome::WorkerLost { .. }
)));
Ok(())
}
#[test]
fn park_disconnected_worker_parks_tasks_without_synthesizing_completions()
-> Result<(), Box<dyn std::error::Error>> {
let (registry, _registration, worker_id) = registry_with_worker()?;
let sink = RecordingSink::default();
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let workflow_id = workflow_id();
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id(60),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id(61),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
let report = tracker.park_disconnected_worker(worker_id, ®istry, &sink)?;
assert_eq!(report.tasks.len(), 2);
assert_eq!(
tracker.in_flight_count()?,
0,
"parking must remove every tracked task so drain accounting reaches zero"
);
assert!(
registry
.workers_for("tenant-a", "default", "charge-card", None)?
.is_empty(),
"the parked worker must be deregistered from routing"
);
let parks = sink
.parks
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
assert_eq!(parks.len(), 2, "each task must be parked exactly once");
drop(parks);
assert!(
sink.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.is_empty(),
"parking must never synthesize an activity completion"
);
let second = tracker.park_disconnected_worker(worker_id, ®istry, &sink)?;
assert!(second.tasks.is_empty());
let third = tracker.fail_disconnected_worker(worker_id, ®istry, &sink)?;
assert!(third.tasks.is_empty());
assert_eq!(
sink.parks
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.len(),
2,
"re-sweeping a parked worker must park nothing further"
);
assert!(
sink.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.is_empty(),
"a fail sweep after the park must fail nothing"
);
Ok(())
}
#[tokio::test]
async fn park_all_in_flight_workers_parks_everything_and_wakes_drain_waiters()
-> Result<(), Box<dyn std::error::Error>> {
let (registry, _registration, worker_id) = registry_with_worker()?;
let sink = RecordingSink::default();
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let workflow_id = workflow_id();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id(70),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
Instant::now(),
)?;
let notified = tracker.empty.notified();
tokio::pin!(notified);
let reports = tracker.park_all_in_flight_workers(®istry, &sink)?;
assert_eq!(reports.len(), 1);
assert_eq!(reports[0].worker_id, worker_id);
assert_eq!(reports[0].tasks.len(), 1);
assert_eq!(tracker.in_flight_count()?, 0);
assert_eq!(
sink.parks
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.len(),
1
);
assert!(
sink.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
.is_empty(),
"the bulk park must never synthesize a completion"
);
tokio::time::timeout(Duration::from_millis(200), notified)
.await
.map_err(|_| "the bulk park must wake drain waiters")?;
Ok(())
}
#[test]
fn payload_free_heartbeat_refreshes_liveness_without_clearing_progress()
-> Result<(), Box<dyn std::error::Error>> {
let window = Duration::from_secs(5);
let tracker = HeartbeatTracker::new(window);
let worker_id = WorkerIdForTest::registered()?;
let workflow_id = workflow_id();
let activity_id = activity_id(12);
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
let progress = payload(&json!({"percent": 80}))?;
tracker.record_heartbeat(
worker_id,
heartbeat(
workflow_id.clone(),
activity_id.clone(),
Some(progress.clone()),
),
start + Duration::from_secs(1),
)?;
let update = tracker.record_heartbeat(
worker_id,
heartbeat(workflow_id.clone(), activity_id.clone(), None),
start + Duration::from_secs(4),
)?;
assert_eq!(
update.liveness.last_progress,
Some(progress),
"a payload-free liveness beat must not erase handler progress"
);
assert!(
tracker.is_live(
worker_id,
&workflow_id,
&activity_id,
start + Duration::from_secs(8)
)?,
"the payload-free beat must still refresh the liveness stamp"
);
Ok(())
}
#[test]
fn malformed_heartbeat_missing_ids_is_wire_error() -> Result<(), Box<dyn std::error::Error>> {
let worker_id = WorkerIdForTest::registered()?;
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let missing = ProtoHeartbeat {
workflow_id: None,
activity_id: Some(ProtoActivityId::from(activity_id(30))),
progress: None,
};
let result = tracker.record_heartbeat(worker_id, missing, Instant::now());
assert!(matches!(result, Err(ServerError::Wire { .. })));
Ok(())
}
#[test]
fn heartbeat_progress_is_not_reported_as_activity_result()
-> Result<(), Box<dyn std::error::Error>> {
let sink = RecordingSink::default();
let worker_id = WorkerIdForTest::registered()?;
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let workflow_id = workflow_id();
let activity_id = activity_id(40);
let now = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
now,
)?;
tracker.record_heartbeat(
worker_id,
heartbeat(
workflow_id,
activity_id,
Some(Payload::new(
ContentType::Json,
b"{\"progress\":1}".to_vec(),
)),
),
now,
)?;
let completions = sink
.completions
.lock()
.map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
assert!(completions.is_empty());
Ok(())
}
struct WorkerIdForTest;
impl WorkerIdForTest {
fn registered() -> Result<WorkerId, ServerError> {
let (_registry, _registration, worker_id) = registry_with_worker()?;
Ok(worker_id)
}
}
#[test]
fn complete_task_reports_whether_the_entry_was_tracked()
-> Result<(), Box<dyn std::error::Error>> {
let tracker = HeartbeatTracker::new(Duration::from_secs(5));
let worker_id = WorkerIdForTest::registered()?;
let workflow_id = workflow_id();
let id = activity_id(50);
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
Instant::now(),
)?;
assert!(tracker.is_tracked(worker_id, &workflow_id, &id)?);
assert!(
tracker.complete_task(worker_id, &workflow_id, &id)?,
"the first completion retires the tracked entry"
);
assert!(!tracker.is_tracked(worker_id, &workflow_id, &id)?);
assert!(
!tracker.complete_task(worker_id, &workflow_id, &id)?,
"a second completion finds nothing to retire"
);
Ok(())
}
#[test]
fn record_liveness_refreshes_stamp_and_ignores_untracked_tasks()
-> Result<(), Box<dyn std::error::Error>> {
let window = Duration::from_secs(5);
let tracker = HeartbeatTracker::new(window);
let worker_id = WorkerIdForTest::registered()?;
let workflow_id = workflow_id();
let id = activity_id(51);
let start = Instant::now();
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: id.clone(),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
start,
)?;
assert!(tracker.record_liveness(worker_id, &workflow_id, &id, start + window)?);
assert!(tracker.is_live(worker_id, &workflow_id, &id, start + window + window)?);
assert!(tracker.expired_workers(start + window + window)?.is_empty());
assert!(!tracker.record_liveness(
worker_id,
&workflow_id,
&activity_id(52),
start + window
)?);
Ok(())
}
#[test]
fn sweep_interval_is_quarter_window_clamped_to_one_second_and_window() {
assert_eq!(
sweep_interval(Duration::from_secs(30)),
Duration::from_millis(7_500)
);
assert_eq!(
sweep_interval(Duration::from_secs(2)),
Duration::from_secs(1)
);
assert_eq!(
sweep_interval(Duration::from_secs(3_600)),
Duration::from_secs(900)
);
assert_eq!(
sweep_interval(Duration::from_millis(200)),
Duration::from_millis(200)
);
assert_eq!(sweep_interval(Duration::ZERO), Duration::from_millis(1));
}
}