Skip to main content

aion_server/worker/
heartbeat.rs

1//! Heartbeat window tracking and lost-worker failure surfacing.
2
3use std::collections::{HashMap, HashSet};
4use std::sync::{Arc, Mutex, MutexGuard};
5use std::time::{Duration, Instant};
6use tokio::sync::Notify;
7
8use aion_core::{ActivityId, Payload, WorkflowId};
9use aion_proto::{ProtoHeartbeat, WireError};
10
11use crate::error::ServerError;
12use crate::worker::dispatch::{
13    ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink, lost_worker_error,
14};
15use crate::worker::registry::{ConnectedWorkerRegistry, WorkerId};
16
17/// In-flight activity assigned to a connected worker.
18#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct InFlightActivity {
20    /// Owning workflow id.
21    pub workflow_id: WorkflowId,
22    /// Correlating activity id.
23    pub activity_id: ActivityId,
24}
25
26/// Observable liveness state for a single in-flight activity.
27#[derive(Clone, Debug, Eq, PartialEq)]
28pub struct TaskLiveness {
29    /// Worker currently responsible for the task.
30    pub worker_id: WorkerId,
31    /// Owning workflow id.
32    pub workflow_id: WorkflowId,
33    /// Correlating activity id.
34    pub activity_id: ActivityId,
35    /// Operator-configured heartbeat window used for expiry checks.
36    pub heartbeat_window: Duration,
37    /// Monotonic timestamp of assignment or the most recent heartbeat.
38    pub last_heartbeat_at: Instant,
39    /// Optional worker progress from the most recent heartbeat.
40    pub last_progress: Option<Payload>,
41}
42
43/// Result of accepting a heartbeat for an in-flight task.
44#[derive(Clone, Debug, Eq, PartialEq)]
45pub struct HeartbeatUpdate {
46    /// Updated liveness after recording the heartbeat.
47    pub liveness: TaskLiveness,
48}
49
50/// Tasks failed because a worker was declared lost.
51#[derive(Clone, Debug, Eq, PartialEq)]
52pub struct LostWorkerReport {
53    /// Lost worker removed from the connected-worker registry.
54    pub worker_id: WorkerId,
55    /// In-flight activities surfaced to the engine as retryable failures.
56    pub tasks: Vec<InFlightActivity>,
57}
58
59#[derive(Clone, Debug, Eq, Hash, PartialEq)]
60struct TaskKey(WorkerId, WorkflowId, ActivityId);
61
62#[derive(Debug, Default)]
63struct HeartbeatState {
64    tasks: HashMap<TaskKey, TaskLiveness>,
65}
66
67/// Per-task liveness tracker for remote-worker streams.
68#[derive(Clone, Debug)]
69pub struct HeartbeatTracker {
70    heartbeat_window: Duration,
71    inner: Arc<Mutex<HeartbeatState>>,
72    empty: Arc<Notify>,
73}
74
75impl HeartbeatTracker {
76    /// Build a tracker using the operator-supplied heartbeat window.
77    #[must_use]
78    pub fn new(heartbeat_window: Duration) -> Self {
79        Self {
80            heartbeat_window,
81            inner: Arc::new(Mutex::new(HeartbeatState::default())),
82            empty: Arc::new(Notify::new()),
83        }
84    }
85
86    /// Track a newly accepted in-flight activity for heartbeat expiry.
87    ///
88    /// # Errors
89    ///
90    /// Returns [`ServerError::LockPoisoned`] if tracker state cannot be trusted.
91    pub fn track_task(
92        &self,
93        worker_id: WorkerId,
94        task: InFlightActivity,
95        now: Instant,
96    ) -> Result<(), ServerError> {
97        let key = TaskKey::new(
98            worker_id,
99            task.workflow_id.clone(),
100            task.activity_id.clone(),
101        );
102        let liveness = TaskLiveness {
103            worker_id,
104            workflow_id: task.workflow_id,
105            activity_id: task.activity_id,
106            heartbeat_window: self.heartbeat_window,
107            last_heartbeat_at: now,
108            last_progress: None,
109        };
110        self.state()?.tasks.insert(key, liveness);
111        Ok(())
112    }
113
114    /// Stop tracking a completed activity and wake drain waiters if this was the last task.
115    ///
116    /// # Errors
117    ///
118    /// Returns [`ServerError::LockPoisoned`] if tracker state cannot be trusted.
119    pub fn complete_task(
120        &self,
121        worker_id: WorkerId,
122        workflow_id: &WorkflowId,
123        activity_id: &ActivityId,
124    ) -> Result<(), ServerError> {
125        let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
126        let became_empty = {
127            let mut state = self.state()?;
128            state.tasks.remove(&key);
129            state.tasks.is_empty()
130        };
131        if became_empty {
132            self.empty.notify_waiters();
133        }
134        Ok(())
135    }
136
137    /// Number of currently tracked in-flight activities.
138    ///
139    /// # Errors
140    ///
141    /// Returns [`ServerError::LockPoisoned`] if tracker state cannot be trusted.
142    pub fn in_flight_count(&self) -> Result<usize, ServerError> {
143        Ok(self.state()?.tasks.len())
144    }
145
146    /// Record a worker heartbeat without completing the activity.
147    ///
148    /// # Errors
149    ///
150    /// Returns a stable wire error for malformed heartbeats or unknown in-flight tasks.
151    pub fn record_heartbeat(
152        &self,
153        worker_id: WorkerId,
154        heartbeat: ProtoHeartbeat,
155        now: Instant,
156    ) -> Result<HeartbeatUpdate, ServerError> {
157        let decoded = DecodedHeartbeat::try_from(heartbeat)?;
158        let key = TaskKey::new(worker_id, decoded.workflow_id, decoded.activity_id);
159        let mut state = self.state()?;
160        let Some(liveness) = state.tasks.get_mut(&key) else {
161            return Err(wire_error("heartbeat task is not in flight"));
162        };
163        liveness.last_heartbeat_at = now;
164        liveness.last_progress = decoded.progress;
165        Ok(HeartbeatUpdate {
166            liveness: liveness.clone(),
167        })
168    }
169
170    /// Return whether an in-flight task is still within its configured heartbeat window.
171    ///
172    /// # Errors
173    ///
174    /// Returns a stable wire error if the task is not tracked, or lock poison if state cannot be trusted.
175    pub fn is_live(
176        &self,
177        worker_id: WorkerId,
178        workflow_id: &WorkflowId,
179        activity_id: &ActivityId,
180        now: Instant,
181    ) -> Result<bool, ServerError> {
182        let key = TaskKey::new(worker_id, workflow_id.clone(), activity_id.clone());
183        let state = self.state()?;
184        let Some(liveness) = state.tasks.get(&key) else {
185            return Err(wire_error("heartbeat task is not in flight"));
186        };
187        Ok(!is_expired(liveness, now))
188    }
189
190    /// Return the workers that have at least one task beyond the configured heartbeat window.
191    ///
192    /// # Errors
193    ///
194    /// Returns [`ServerError::LockPoisoned`] if tracker state cannot be trusted.
195    pub fn expired_workers(&self, now: Instant) -> Result<Vec<WorkerId>, ServerError> {
196        let state = self.state()?;
197        let mut seen = HashSet::new();
198        let mut workers = Vec::new();
199        for liveness in state.tasks.values() {
200            if is_expired(liveness, now) && seen.insert(liveness.worker_id) {
201                workers.push(liveness.worker_id);
202            }
203        }
204        workers.sort_unstable();
205        Ok(workers)
206    }
207
208    /// Mark all currently expired workers lost and fail their in-flight tasks through the engine sink.
209    ///
210    /// # Errors
211    ///
212    /// Returns registry, tracker, or sink errors without retrying or rescheduling activities.
213    pub fn fail_expired_workers(
214        &self,
215        registry: &ConnectedWorkerRegistry,
216        sink: &impl ActivityCompletionSink,
217        now: Instant,
218    ) -> Result<Vec<LostWorkerReport>, ServerError> {
219        let mut reports = Vec::new();
220        for worker_id in self.expired_workers(now)? {
221            let report = self.fail_lost_worker(worker_id, registry, sink)?;
222            if !report.tasks.is_empty() {
223                reports.push(report);
224            }
225        }
226        Ok(reports)
227    }
228
229    /// Mark a disconnected worker lost and fail its in-flight tasks through the engine sink.
230    ///
231    /// # Errors
232    ///
233    /// Returns registry, tracker, or sink errors without retrying or rescheduling activities.
234    pub fn fail_disconnected_worker(
235        &self,
236        worker_id: WorkerId,
237        registry: &ConnectedWorkerRegistry,
238        sink: &impl ActivityCompletionSink,
239    ) -> Result<LostWorkerReport, ServerError> {
240        self.fail_lost_worker(worker_id, registry, sink)
241    }
242
243    /// Mark every currently in-flight worker lost and fail all remaining tasks through the sink.
244    ///
245    /// # Errors
246    ///
247    /// Returns registry, tracker, or sink errors without retrying or rescheduling activities.
248    pub fn fail_all_in_flight_workers(
249        &self,
250        registry: &ConnectedWorkerRegistry,
251        sink: &impl ActivityCompletionSink,
252    ) -> Result<Vec<LostWorkerReport>, ServerError> {
253        let worker_ids = {
254            let state = self.state()?;
255            let mut worker_ids = state
256                .tasks
257                .values()
258                .map(|liveness| liveness.worker_id)
259                .collect::<HashSet<_>>()
260                .into_iter()
261                .collect::<Vec<_>>();
262            worker_ids.sort_unstable();
263            worker_ids
264        };
265        let mut reports = Vec::new();
266        for worker_id in worker_ids {
267            let report = self.fail_lost_worker(worker_id, registry, sink)?;
268            if !report.tasks.is_empty() {
269                reports.push(report);
270            }
271        }
272        self.empty.notify_waiters();
273        Ok(reports)
274    }
275
276    fn fail_lost_worker(
277        &self,
278        worker_id: WorkerId,
279        registry: &ConnectedWorkerRegistry,
280        sink: &impl ActivityCompletionSink,
281    ) -> Result<LostWorkerReport, ServerError> {
282        // Deregister BEFORE collecting tasks: the dispatch path tracks its
283        // task, sends, and then checks `registry.is_registered`. With this
284        // ordering, a dispatch that still sees the worker registered is
285        // guaranteed its tracked task is visible to any later sweep, so the
286        // unbounded completion wait always gets a lost-worker failure. (The
287        // reverse order leaves a window where a task tracked between the
288        // collection and the deregistration is never failed by anyone.)
289        registry.deregister(worker_id)?;
290        let tasks = self.remove_worker_tasks(worker_id)?;
291        for task in &tasks {
292            sink.complete_activity(ActivityCompletion {
293                workflow_id: task.workflow_id.clone(),
294                activity_id: task.activity_id.clone(),
295                outcome: ActivityCompletionOutcome::Failed(lost_worker_error(worker_id)),
296            })?;
297        }
298        Ok(LostWorkerReport { worker_id, tasks })
299    }
300
301    fn remove_worker_tasks(
302        &self,
303        worker_id: WorkerId,
304    ) -> Result<Vec<InFlightActivity>, ServerError> {
305        let mut state = self.state()?;
306        let keys = state
307            .tasks
308            .keys()
309            .filter(|key| key.worker_id() == worker_id)
310            .cloned()
311            .collect::<Vec<_>>();
312        let mut tasks = Vec::with_capacity(keys.len());
313        for key in keys {
314            if let Some(liveness) = state.tasks.remove(&key) {
315                tasks.push(InFlightActivity {
316                    workflow_id: liveness.workflow_id,
317                    activity_id: liveness.activity_id,
318                });
319            }
320        }
321        Ok(tasks)
322    }
323
324    fn state(&self) -> Result<MutexGuard<'_, HeartbeatState>, ServerError> {
325        self.inner
326            .lock()
327            .map_err(|_| ServerError::lock_poisoned("worker heartbeat tracker"))
328    }
329}
330
331impl TaskKey {
332    fn new(worker_id: WorkerId, workflow_id: WorkflowId, activity_id: ActivityId) -> Self {
333        Self(worker_id, workflow_id, activity_id)
334    }
335
336    const fn worker_id(&self) -> WorkerId {
337        self.0
338    }
339}
340
341struct DecodedHeartbeat {
342    workflow_id: WorkflowId,
343    activity_id: ActivityId,
344    progress: Option<Payload>,
345}
346
347impl TryFrom<ProtoHeartbeat> for DecodedHeartbeat {
348    type Error = ServerError;
349
350    fn try_from(value: ProtoHeartbeat) -> Result<Self, Self::Error> {
351        let workflow_id = value
352            .workflow_id
353            .ok_or_else(|| wire_error("heartbeat workflow id is missing"))
354            .and_then(|id| WorkflowId::try_from(id).map_err(ServerError::from))?;
355        let activity_id = value
356            .activity_id
357            .ok_or_else(|| wire_error("heartbeat activity id is missing"))
358            .map(ActivityId::from)?;
359        let progress = value
360            .progress
361            .map(Payload::try_from)
362            .transpose()
363            .map_err(ServerError::from)?;
364        Ok(Self {
365            workflow_id,
366            activity_id,
367            progress,
368        })
369    }
370}
371
372fn is_expired(liveness: &TaskLiveness, now: Instant) -> bool {
373    now.checked_duration_since(liveness.last_heartbeat_at)
374        .is_some_and(|elapsed| elapsed > liveness.heartbeat_window)
375}
376
377fn wire_error(message: &'static str) -> ServerError {
378    ServerError::Wire {
379        wire: WireError::backend(message),
380    }
381}
382
383#[cfg(test)]
384mod tests {
385    use std::sync::Mutex;
386
387    use aion_core::{ActivityErrorKind, ContentType};
388    use aion_proto::{ProtoActivityId, ProtoPayload, ProtoWorkflowId};
389    use serde_json::json;
390    use uuid::Uuid;
391
392    use crate::worker::registry::WorkerRegistration;
393
394    use super::*;
395
396    #[derive(Default)]
397    struct RecordingSink {
398        completions: Mutex<Vec<ActivityCompletion>>,
399    }
400
401    impl ActivityCompletionSink for RecordingSink {
402        fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
403            self.completions
404                .lock()
405                .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
406                .push(completion);
407            Ok(())
408        }
409    }
410
411    fn workflow_id() -> WorkflowId {
412        WorkflowId::new(Uuid::nil())
413    }
414
415    fn activity_id(position: u64) -> ActivityId {
416        ActivityId::from_sequence_position(position)
417    }
418
419    fn payload(value: &serde_json::Value) -> Result<Payload, Box<dyn std::error::Error>> {
420        Ok(Payload::from_json(value)?)
421    }
422
423    fn heartbeat(
424        workflow_id: WorkflowId,
425        activity_id: ActivityId,
426        progress: Option<Payload>,
427    ) -> ProtoHeartbeat {
428        ProtoHeartbeat {
429            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
430            activity_id: Some(ProtoActivityId::from(activity_id)),
431            progress: progress.map(ProtoPayload::from),
432        }
433    }
434
435    fn registry_with_worker()
436    -> Result<(ConnectedWorkerRegistry, WorkerRegistration, WorkerId), ServerError> {
437        let registry = ConnectedWorkerRegistry::default();
438        let (tx, _rx) = tokio::sync::mpsc::channel(1);
439        let activity_types = [String::from("charge-card")];
440        let registration = registry.register("tenant-a", activity_types.iter(), tx)?;
441        let worker_id = registration
442            .worker_id()
443            .ok_or_else(|| ServerError::lock_poisoned("test worker registration"))?;
444        Ok((registry, registration, worker_id))
445    }
446
447    #[test]
448    fn heartbeat_refresh_keeps_task_live_across_window() -> Result<(), Box<dyn std::error::Error>> {
449        let window = Duration::from_secs(5);
450        let tracker = HeartbeatTracker::new(window);
451        let worker_id = WorkerIdForTest::registered()?;
452        let workflow_id = workflow_id();
453        let activity_id = activity_id(10);
454        let start = Instant::now();
455
456        tracker.track_task(
457            worker_id,
458            InFlightActivity {
459                workflow_id: workflow_id.clone(),
460                activity_id: activity_id.clone(),
461            },
462            start,
463        )?;
464        assert!(tracker.is_live(worker_id, &workflow_id, &activity_id, start + window)?);
465
466        let progress = payload(&json!({"percent": 50}))?;
467        let update = tracker.record_heartbeat(
468            worker_id,
469            heartbeat(
470                workflow_id.clone(),
471                activity_id.clone(),
472                Some(progress.clone()),
473            ),
474            start + window,
475        )?;
476
477        assert_eq!(update.liveness.last_progress, Some(progress));
478        assert!(tracker.is_live(
479            worker_id,
480            &workflow_id,
481            &activity_id,
482            start + window + window
483        )?);
484        assert!(tracker.expired_workers(start + window + window)?.is_empty());
485        Ok(())
486    }
487
488    #[test]
489    fn missed_heartbeat_deregisters_worker_and_fails_in_flight_once()
490    -> Result<(), Box<dyn std::error::Error>> {
491        let (registry, _registration, worker_id) = registry_with_worker()?;
492        let sink = RecordingSink::default();
493        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
494        let workflow_id = workflow_id();
495        let activity_id = activity_id(11);
496        let start = Instant::now();
497
498        tracker.track_task(
499            worker_id,
500            InFlightActivity {
501                workflow_id: workflow_id.clone(),
502                activity_id: activity_id.clone(),
503            },
504            start,
505        )?;
506
507        let reports =
508            tracker.fail_expired_workers(&registry, &sink, start + Duration::from_secs(6))?;
509        assert_eq!(reports.len(), 1);
510        assert_eq!(reports[0].worker_id, worker_id);
511        assert_eq!(reports[0].tasks.len(), 1);
512        assert!(registry.workers_for("tenant-a", "charge-card")?.is_empty());
513
514        let second = tracker.fail_disconnected_worker(worker_id, &registry, &sink)?;
515        assert!(second.tasks.is_empty());
516        let completions = sink
517            .completions
518            .lock()
519            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
520        assert_eq!(completions.len(), 1);
521        assert_eq!(completions[0].workflow_id, workflow_id);
522        assert_eq!(completions[0].activity_id, activity_id);
523        match &completions[0].outcome {
524            ActivityCompletionOutcome::Failed(error) => {
525                assert_eq!(error.kind, ActivityErrorKind::Retryable);
526                assert!(error.is_retryable());
527            }
528            ActivityCompletionOutcome::Succeeded(_) => {
529                return Err("expected lost-worker failure".into());
530            }
531        }
532        Ok(())
533    }
534
535    #[test]
536    fn disconnected_worker_fails_each_in_flight_task_once() -> Result<(), Box<dyn std::error::Error>>
537    {
538        let (registry, _registration, worker_id) = registry_with_worker()?;
539        let sink = RecordingSink::default();
540        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
541        let workflow_id = workflow_id();
542        let start = Instant::now();
543
544        tracker.track_task(
545            worker_id,
546            InFlightActivity {
547                workflow_id: workflow_id.clone(),
548                activity_id: activity_id(21),
549            },
550            start,
551        )?;
552        tracker.track_task(
553            worker_id,
554            InFlightActivity {
555                workflow_id,
556                activity_id: activity_id(22),
557            },
558            start,
559        )?;
560
561        let report = tracker.fail_disconnected_worker(worker_id, &registry, &sink)?;
562        assert_eq!(report.tasks.len(), 2);
563        assert!(registry.workers_for("tenant-a", "charge-card")?.is_empty());
564
565        let completions = sink
566            .completions
567            .lock()
568            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
569        assert_eq!(completions.len(), 2);
570        assert!(completions.iter().all(|completion| matches!(
571            &completion.outcome,
572            ActivityCompletionOutcome::Failed(error)
573                if error.kind == ActivityErrorKind::Retryable && error.is_retryable()
574        )));
575        Ok(())
576    }
577
578    #[test]
579    fn malformed_heartbeat_missing_ids_is_wire_error() -> Result<(), Box<dyn std::error::Error>> {
580        let worker_id = WorkerIdForTest::registered()?;
581        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
582        let missing = ProtoHeartbeat {
583            workflow_id: None,
584            activity_id: Some(ProtoActivityId::from(activity_id(30))),
585            progress: None,
586        };
587
588        let result = tracker.record_heartbeat(worker_id, missing, Instant::now());
589        assert!(matches!(result, Err(ServerError::Wire { .. })));
590        Ok(())
591    }
592
593    #[test]
594    fn heartbeat_progress_is_not_reported_as_activity_result()
595    -> Result<(), Box<dyn std::error::Error>> {
596        let sink = RecordingSink::default();
597        let worker_id = WorkerIdForTest::registered()?;
598        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
599        let workflow_id = workflow_id();
600        let activity_id = activity_id(40);
601        let now = Instant::now();
602
603        tracker.track_task(
604            worker_id,
605            InFlightActivity {
606                workflow_id: workflow_id.clone(),
607                activity_id: activity_id.clone(),
608            },
609            now,
610        )?;
611        tracker.record_heartbeat(
612            worker_id,
613            heartbeat(
614                workflow_id,
615                activity_id,
616                Some(Payload::new(
617                    ContentType::Json,
618                    b"{\"progress\":1}".to_vec(),
619                )),
620            ),
621            now,
622        )?;
623
624        let completions = sink
625            .completions
626            .lock()
627            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
628        assert!(completions.is_empty());
629        Ok(())
630    }
631
632    struct WorkerIdForTest;
633
634    impl WorkerIdForTest {
635        fn registered() -> Result<WorkerId, ServerError> {
636            let (_registry, _registration, worker_id) = registry_with_worker()?;
637            Ok(worker_id)
638        }
639    }
640}