1use 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#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct InFlightActivity {
20 pub workflow_id: WorkflowId,
22 pub activity_id: ActivityId,
24}
25
26#[derive(Clone, Debug, Eq, PartialEq)]
28pub struct TaskLiveness {
29 pub worker_id: WorkerId,
31 pub workflow_id: WorkflowId,
33 pub activity_id: ActivityId,
35 pub heartbeat_window: Duration,
37 pub last_heartbeat_at: Instant,
39 pub last_progress: Option<Payload>,
41}
42
43#[derive(Clone, Debug, Eq, PartialEq)]
45pub struct HeartbeatUpdate {
46 pub liveness: TaskLiveness,
48}
49
50#[derive(Clone, Debug, Eq, PartialEq)]
52pub struct LostWorkerReport {
53 pub worker_id: WorkerId,
55 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#[derive(Clone, Debug)]
69pub struct HeartbeatTracker {
70 heartbeat_window: Duration,
71 inner: Arc<Mutex<HeartbeatState>>,
72 empty: Arc<Notify>,
73}
74
75impl HeartbeatTracker {
76 #[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 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 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 pub fn in_flight_count(&self) -> Result<usize, ServerError> {
143 Ok(self.state()?.tasks.len())
144 }
145
146 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 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 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 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 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 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 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(®istry, &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, ®istry, &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, ®istry, &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}