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                run_id: None,
296                outcome: ActivityCompletionOutcome::Failed(lost_worker_error(worker_id)),
297            })?;
298        }
299        Ok(LostWorkerReport { worker_id, tasks })
300    }
301
302    fn remove_worker_tasks(
303        &self,
304        worker_id: WorkerId,
305    ) -> Result<Vec<InFlightActivity>, ServerError> {
306        let mut state = self.state()?;
307        let keys = state
308            .tasks
309            .keys()
310            .filter(|key| key.worker_id() == worker_id)
311            .cloned()
312            .collect::<Vec<_>>();
313        let mut tasks = Vec::with_capacity(keys.len());
314        for key in keys {
315            if let Some(liveness) = state.tasks.remove(&key) {
316                tasks.push(InFlightActivity {
317                    workflow_id: liveness.workflow_id,
318                    activity_id: liveness.activity_id,
319                });
320            }
321        }
322        Ok(tasks)
323    }
324
325    fn state(&self) -> Result<MutexGuard<'_, HeartbeatState>, ServerError> {
326        self.inner
327            .lock()
328            .map_err(|_| ServerError::lock_poisoned("worker heartbeat tracker"))
329    }
330}
331
332impl TaskKey {
333    fn new(worker_id: WorkerId, workflow_id: WorkflowId, activity_id: ActivityId) -> Self {
334        Self(worker_id, workflow_id, activity_id)
335    }
336
337    const fn worker_id(&self) -> WorkerId {
338        self.0
339    }
340}
341
342struct DecodedHeartbeat {
343    workflow_id: WorkflowId,
344    activity_id: ActivityId,
345    progress: Option<Payload>,
346}
347
348impl TryFrom<ProtoHeartbeat> for DecodedHeartbeat {
349    type Error = ServerError;
350
351    fn try_from(value: ProtoHeartbeat) -> Result<Self, Self::Error> {
352        let workflow_id = value
353            .workflow_id
354            .ok_or_else(|| wire_error("heartbeat workflow id is missing"))
355            .and_then(|id| WorkflowId::try_from(id).map_err(ServerError::from))?;
356        let activity_id = value
357            .activity_id
358            .ok_or_else(|| wire_error("heartbeat activity id is missing"))
359            .map(ActivityId::from)?;
360        let progress = value
361            .progress
362            .map(Payload::try_from)
363            .transpose()
364            .map_err(ServerError::from)?;
365        Ok(Self {
366            workflow_id,
367            activity_id,
368            progress,
369        })
370    }
371}
372
373fn is_expired(liveness: &TaskLiveness, now: Instant) -> bool {
374    now.checked_duration_since(liveness.last_heartbeat_at)
375        .is_some_and(|elapsed| elapsed > liveness.heartbeat_window)
376}
377
378fn wire_error(message: &'static str) -> ServerError {
379    ServerError::Wire {
380        wire: WireError::backend(message),
381    }
382}
383
384#[cfg(test)]
385mod tests {
386    use std::sync::Mutex;
387
388    use aion_core::{ActivityErrorKind, ContentType};
389    use aion_proto::{ProtoActivityId, ProtoPayload, ProtoWorkflowId};
390    use serde_json::json;
391    use uuid::Uuid;
392
393    use crate::worker::registry::WorkerRegistration;
394
395    use super::*;
396
397    #[derive(Default)]
398    struct RecordingSink {
399        completions: Mutex<Vec<ActivityCompletion>>,
400    }
401
402    impl ActivityCompletionSink for RecordingSink {
403        fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
404            self.completions
405                .lock()
406                .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
407                .push(completion);
408            Ok(())
409        }
410    }
411
412    fn workflow_id() -> WorkflowId {
413        WorkflowId::new(Uuid::nil())
414    }
415
416    fn activity_id(position: u64) -> ActivityId {
417        ActivityId::from_sequence_position(position)
418    }
419
420    fn payload(value: &serde_json::Value) -> Result<Payload, Box<dyn std::error::Error>> {
421        Ok(Payload::from_json(value)?)
422    }
423
424    fn heartbeat(
425        workflow_id: WorkflowId,
426        activity_id: ActivityId,
427        progress: Option<Payload>,
428    ) -> ProtoHeartbeat {
429        ProtoHeartbeat {
430            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
431            activity_id: Some(ProtoActivityId::from(activity_id)),
432            progress: progress.map(ProtoPayload::from),
433        }
434    }
435
436    fn registry_with_worker()
437    -> Result<(ConnectedWorkerRegistry, WorkerRegistration, WorkerId), ServerError> {
438        let registry = ConnectedWorkerRegistry::default();
439        let (tx, _rx) = tokio::sync::mpsc::channel(1);
440        let activity_types = [String::from("charge-card")];
441        let registration = registry.register("tenant-a", activity_types.iter(), tx)?;
442        let worker_id = registration
443            .worker_id()
444            .ok_or_else(|| ServerError::lock_poisoned("test worker registration"))?;
445        Ok((registry, registration, worker_id))
446    }
447
448    #[test]
449    fn heartbeat_refresh_keeps_task_live_across_window() -> Result<(), Box<dyn std::error::Error>> {
450        let window = Duration::from_secs(5);
451        let tracker = HeartbeatTracker::new(window);
452        let worker_id = WorkerIdForTest::registered()?;
453        let workflow_id = workflow_id();
454        let activity_id = activity_id(10);
455        let start = Instant::now();
456
457        tracker.track_task(
458            worker_id,
459            InFlightActivity {
460                workflow_id: workflow_id.clone(),
461                activity_id: activity_id.clone(),
462            },
463            start,
464        )?;
465        assert!(tracker.is_live(worker_id, &workflow_id, &activity_id, start + window)?);
466
467        let progress = payload(&json!({"percent": 50}))?;
468        let update = tracker.record_heartbeat(
469            worker_id,
470            heartbeat(
471                workflow_id.clone(),
472                activity_id.clone(),
473                Some(progress.clone()),
474            ),
475            start + window,
476        )?;
477
478        assert_eq!(update.liveness.last_progress, Some(progress));
479        assert!(tracker.is_live(
480            worker_id,
481            &workflow_id,
482            &activity_id,
483            start + window + window
484        )?);
485        assert!(tracker.expired_workers(start + window + window)?.is_empty());
486        Ok(())
487    }
488
489    #[test]
490    fn missed_heartbeat_deregisters_worker_and_fails_in_flight_once()
491    -> Result<(), Box<dyn std::error::Error>> {
492        let (registry, _registration, worker_id) = registry_with_worker()?;
493        let sink = RecordingSink::default();
494        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
495        let workflow_id = workflow_id();
496        let activity_id = activity_id(11);
497        let start = Instant::now();
498
499        tracker.track_task(
500            worker_id,
501            InFlightActivity {
502                workflow_id: workflow_id.clone(),
503                activity_id: activity_id.clone(),
504            },
505            start,
506        )?;
507
508        let reports =
509            tracker.fail_expired_workers(&registry, &sink, start + Duration::from_secs(6))?;
510        assert_eq!(reports.len(), 1);
511        assert_eq!(reports[0].worker_id, worker_id);
512        assert_eq!(reports[0].tasks.len(), 1);
513        assert!(
514            registry
515                .workers_for("tenant-a", "default", "charge-card", None)?
516                .is_empty()
517        );
518
519        let second = tracker.fail_disconnected_worker(worker_id, &registry, &sink)?;
520        assert!(second.tasks.is_empty());
521        let completions = sink
522            .completions
523            .lock()
524            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
525        assert_eq!(completions.len(), 1);
526        assert_eq!(completions[0].workflow_id, workflow_id);
527        assert_eq!(completions[0].activity_id, activity_id);
528        match &completions[0].outcome {
529            ActivityCompletionOutcome::Failed(error) => {
530                assert_eq!(error.kind, ActivityErrorKind::Retryable);
531                assert!(error.is_retryable());
532            }
533            ActivityCompletionOutcome::Succeeded(_) => {
534                return Err("expected lost-worker failure".into());
535            }
536        }
537        Ok(())
538    }
539
540    #[test]
541    fn disconnected_worker_fails_each_in_flight_task_once() -> Result<(), Box<dyn std::error::Error>>
542    {
543        let (registry, _registration, worker_id) = registry_with_worker()?;
544        let sink = RecordingSink::default();
545        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
546        let workflow_id = workflow_id();
547        let start = Instant::now();
548
549        tracker.track_task(
550            worker_id,
551            InFlightActivity {
552                workflow_id: workflow_id.clone(),
553                activity_id: activity_id(21),
554            },
555            start,
556        )?;
557        tracker.track_task(
558            worker_id,
559            InFlightActivity {
560                workflow_id,
561                activity_id: activity_id(22),
562            },
563            start,
564        )?;
565
566        let report = tracker.fail_disconnected_worker(worker_id, &registry, &sink)?;
567        assert_eq!(report.tasks.len(), 2);
568        assert!(
569            registry
570                .workers_for("tenant-a", "default", "charge-card", None)?
571                .is_empty()
572        );
573
574        let completions = sink
575            .completions
576            .lock()
577            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
578        assert_eq!(completions.len(), 2);
579        assert!(completions.iter().all(|completion| matches!(
580            &completion.outcome,
581            ActivityCompletionOutcome::Failed(error)
582                if error.kind == ActivityErrorKind::Retryable && error.is_retryable()
583        )));
584        Ok(())
585    }
586
587    #[test]
588    fn malformed_heartbeat_missing_ids_is_wire_error() -> Result<(), Box<dyn std::error::Error>> {
589        let worker_id = WorkerIdForTest::registered()?;
590        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
591        let missing = ProtoHeartbeat {
592            workflow_id: None,
593            activity_id: Some(ProtoActivityId::from(activity_id(30))),
594            progress: None,
595        };
596
597        let result = tracker.record_heartbeat(worker_id, missing, Instant::now());
598        assert!(matches!(result, Err(ServerError::Wire { .. })));
599        Ok(())
600    }
601
602    #[test]
603    fn heartbeat_progress_is_not_reported_as_activity_result()
604    -> Result<(), Box<dyn std::error::Error>> {
605        let sink = RecordingSink::default();
606        let worker_id = WorkerIdForTest::registered()?;
607        let tracker = HeartbeatTracker::new(Duration::from_secs(5));
608        let workflow_id = workflow_id();
609        let activity_id = activity_id(40);
610        let now = Instant::now();
611
612        tracker.track_task(
613            worker_id,
614            InFlightActivity {
615                workflow_id: workflow_id.clone(),
616                activity_id: activity_id.clone(),
617            },
618            now,
619        )?;
620        tracker.record_heartbeat(
621            worker_id,
622            heartbeat(
623                workflow_id,
624                activity_id,
625                Some(Payload::new(
626                    ContentType::Json,
627                    b"{\"progress\":1}".to_vec(),
628                )),
629            ),
630            now,
631        )?;
632
633        let completions = sink
634            .completions
635            .lock()
636            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
637        assert!(completions.is_empty());
638        Ok(())
639    }
640
641    struct WorkerIdForTest;
642
643    impl WorkerIdForTest {
644        fn registered() -> Result<WorkerId, ServerError> {
645            let (_registry, _registration, worker_id) = registry_with_worker()?;
646            Ok(worker_id)
647        }
648    }
649}