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 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(®istry, &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, ®istry, &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, ®istry, &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}