Skip to main content

aion_worker/protocol/
reconnect.rs

1//! Backoff reconnect, re-register, and re-report un-acked results.
2
3use std::collections::{BTreeMap, BTreeSet};
4use std::future::Future;
5use std::time::Duration;
6
7use aion_core::{ActivityError, ActivityId, Payload, RunId, WorkflowId};
8use tracing::{debug, error, warn};
9use uuid::Uuid;
10
11use crate::config::WorkerConfig;
12use crate::error::WorkerError;
13use crate::protocol::{GrpcWorkerSession, WorkerSession};
14
15/// Result or failure computed locally and not yet acknowledged by the engine.
16#[derive(Clone, Debug, PartialEq, Eq)]
17pub enum PendingActivityReport {
18    /// Successful activity output to re-report until acknowledged.
19    Completed {
20        /// Workflow owning the activity.
21        workflow_id: WorkflowId,
22        /// Activity identifier used by AW for idempotent ingest.
23        activity_id: ActivityId,
24        /// Concrete workflow run to echo on re-report, when known.
25        run_id: Option<RunId>,
26        /// Opaque execution generation to echo on every re-report.
27        completion_token: String,
28        /// Opaque activity output payload.
29        output: Payload,
30    },
31    /// Explicitly classified activity failure to re-report until acknowledged.
32    Failed {
33        /// Workflow owning the activity.
34        workflow_id: WorkflowId,
35        /// Activity identifier used by AW for idempotent ingest.
36        activity_id: ActivityId,
37        /// Concrete workflow run to echo on re-report, when known.
38        run_id: Option<RunId>,
39        /// Opaque execution generation to echo on every re-report.
40        completion_token: String,
41        /// Classified activity error.
42        failure: ActivityError,
43    },
44}
45
46impl PendingActivityReport {
47    /// Returns the report's activity id key.
48    #[must_use]
49    pub const fn activity_id(&self) -> &ActivityId {
50        match self {
51            Self::Completed { activity_id, .. } | Self::Failed { activity_id, .. } => activity_id,
52        }
53    }
54
55    /// Returns the workflow owning the report's activity.
56    #[must_use]
57    pub const fn workflow_id(&self) -> &WorkflowId {
58        match self {
59            Self::Completed { workflow_id, .. } | Self::Failed { workflow_id, .. } => workflow_id,
60        }
61    }
62}
63
64/// Deterministic tracker key: activity ids are sequence positions scoped to
65/// one workflow, so distinct workflows legitimately collide on the bare
66/// position and must be keyed by workflow as well.
67type PendingReportKey = (Uuid, u64);
68
69fn pending_report_key(workflow_id: &WorkflowId, activity_id: &ActivityId) -> PendingReportKey {
70    (workflow_id.as_uuid(), activity_id.sequence_position())
71}
72
73/// In-memory source of truth for locally reported results awaiting engine ack.
74#[derive(Clone, Debug, Default, PartialEq, Eq)]
75pub struct UnackedResultTracker {
76    reports: BTreeMap<PendingReportKey, PendingActivityReport>,
77}
78
79impl UnackedResultTracker {
80    /// Creates an empty tracker.
81    #[must_use]
82    pub const fn new() -> Self {
83        Self {
84            reports: BTreeMap::new(),
85        }
86    }
87
88    /// Records a report, replacing any earlier pending report for the same
89    /// workflow and activity id.
90    pub fn record(&mut self, report: PendingActivityReport) {
91        let key = pending_report_key(report.workflow_id(), report.activity_id());
92        self.reports.insert(key, report);
93    }
94
95    /// Drops a report once the engine explicitly acknowledges it.
96    pub fn acknowledge(
97        &mut self,
98        workflow_id: &WorkflowId,
99        activity_id: &ActivityId,
100    ) -> Option<PendingActivityReport> {
101        self.reports
102            .remove(&pending_report_key(workflow_id, activity_id))
103    }
104
105    /// Returns the number of unacknowledged reports.
106    #[must_use]
107    pub fn len(&self) -> usize {
108        self.reports.len()
109    }
110
111    /// Returns true when no reports are waiting for acknowledgement.
112    #[must_use]
113    pub fn is_empty(&self) -> bool {
114        self.reports.is_empty()
115    }
116
117    /// Gets a pending report by its workflow and activity id.
118    #[must_use]
119    pub fn get(
120        &self,
121        workflow_id: &WorkflowId,
122        activity_id: &ActivityId,
123    ) -> Option<&PendingActivityReport> {
124        self.reports
125            .get(&pending_report_key(workflow_id, activity_id))
126    }
127
128    /// Returns a deterministic snapshot for re-reporting without holding a borrow.
129    #[must_use]
130    pub fn snapshot(&self) -> Vec<PendingActivityReport> {
131        self.reports.values().cloned().collect()
132    }
133}
134
135/// Validated reconnect backoff settings drawn from [`WorkerConfig`].
136#[derive(Clone, Debug, PartialEq, Eq)]
137pub struct ReconnectBackoff {
138    initial: Duration,
139    max: Duration,
140    attempts: usize,
141}
142
143impl ReconnectBackoff {
144    /// Builds reconnect backoff from worker config.
145    ///
146    /// # Errors
147    ///
148    /// Returns [`WorkerError::Registration`] if delays or attempt counts are zero.
149    pub fn from_config(config: &WorkerConfig) -> Result<Self, WorkerError> {
150        if config.reconnect.initial_backoff.is_zero() {
151            return Err(WorkerError::registration(InvalidReconnectBackoff {
152                message: String::from("reconnect initial_backoff must be greater than zero"),
153            }));
154        }
155        if config.reconnect.max_backoff.is_zero() {
156            return Err(WorkerError::registration(InvalidReconnectBackoff {
157                message: String::from("reconnect max_backoff must be greater than zero"),
158            }));
159        }
160        if config.reconnect.max_attempts == 0 {
161            return Err(WorkerError::registration(InvalidReconnectBackoff {
162                message: String::from("reconnect max_attempts must be greater than zero"),
163            }));
164        }
165        Ok(Self {
166            initial: config.reconnect.initial_backoff,
167            max: config.reconnect.max_backoff,
168            attempts: config.reconnect.max_attempts,
169        })
170    }
171
172    /// Returns the bounded exponential delay after `completed_failures` failures.
173    ///
174    /// The delay doubles per completed failure starting from the configured
175    /// initial backoff and is capped at the configured maximum backoff.
176    #[must_use]
177    pub fn delay_for_attempt(&self, completed_failures: usize) -> Duration {
178        let bounded_shift = completed_failures.saturating_sub(1).min(31);
179        let shift = u32::try_from(bounded_shift).map_or(31, |shift| shift);
180        let factor = 1_u32.checked_shl(shift).map_or(u32::MAX, |factor| factor);
181        self.initial.saturating_mul(factor).min(self.max)
182    }
183
184    /// Returns the configured maximum number of reconnect attempts.
185    #[must_use]
186    pub const fn attempts(&self) -> usize {
187        self.attempts
188    }
189
190    /// Returns the configured maximum backoff delay cap.
191    ///
192    /// The run loop also uses this as its session-health threshold: the cap
193    /// is the policy's own definition of the longest pause, so an
194    /// established session that survives longer than it is demonstrably past
195    /// the flapping regime and resets the cumulative drop budget when it
196    /// eventually drops.
197    #[must_use]
198    pub const fn max_delay(&self) -> Duration {
199        self.max
200    }
201}
202
203/// Connects, handshakes, and registers a fresh gRPC worker session.
204///
205/// # Errors
206///
207/// Returns [`WorkerError`] if connection, handshake, or registration fails.
208pub async fn connect_registered_grpc_session(
209    config: &WorkerConfig,
210    activity_types: Vec<String>,
211    activities: Vec<aion_package::ActivityDescriptor>,
212    available_handlers: &BTreeSet<String>,
213) -> Result<GrpcWorkerSession, WorkerError> {
214    let session = GrpcWorkerSession::connect(config.clone()).await?;
215    register_connected_session(
216        session,
217        config,
218        activity_types,
219        activities,
220        available_handlers,
221    )
222    .await
223}
224
225/// Handshakes and registers an already-connected session.
226///
227/// # Errors
228///
229/// Returns [`WorkerError`] if handshake or registration fails.
230pub async fn register_connected_session<S>(
231    mut session: S,
232    config: &WorkerConfig,
233    activity_types: Vec<String>,
234    activities: Vec<aion_package::ActivityDescriptor>,
235    available_handlers: &BTreeSet<String>,
236) -> Result<S, WorkerError>
237where
238    S: WorkerSession,
239{
240    session.handshake(config).await?;
241    session
242        .register_with_contract(activity_types, activities, available_handlers)
243        .await?;
244    Ok(session)
245}
246
247/// Reconnects with bounded exponential backoff using an injected session factory.
248///
249/// # Errors
250///
251/// Returns the last [`WorkerError`] after configured attempts are exhausted, or
252/// immediately when the failure is a non-retryable `PermissionDenied` /
253/// `Unauthenticated` denial, or if the config contains invalid zero reconnect
254/// settings.
255pub async fn reconnect_with_backoff<S, F, Fut>(
256    config: &WorkerConfig,
257    activity_types: Vec<String>,
258    activities: Vec<aion_package::ActivityDescriptor>,
259    available_handlers: &BTreeSet<String>,
260    connect: F,
261) -> Result<S, WorkerError>
262where
263    S: WorkerSession,
264    F: FnMut() -> Fut,
265    Fut: Future<Output = Result<S, WorkerError>>,
266{
267    reconnect_with_sleep(
268        config,
269        activity_types,
270        activities,
271        available_handlers,
272        connect,
273        tokio::time::sleep,
274    )
275    .await
276}
277
278/// Testable reconnect helper with injectable sleep.
279///
280/// # Errors
281///
282/// Returns the last [`WorkerError`] after configured attempts are exhausted, or
283/// immediately — without consuming further attempts — when a failure is a
284/// non-retryable denial ([`WorkerError::is_retryable`] is false, i.e. the
285/// server answered `PermissionDenied` or `Unauthenticated`), or if the config
286/// contains invalid zero reconnect settings.
287pub async fn reconnect_with_sleep<S, F, Fut, Sleep, SleepFut>(
288    config: &WorkerConfig,
289    activity_types: Vec<String>,
290    activities: Vec<aion_package::ActivityDescriptor>,
291    available_handlers: &BTreeSet<String>,
292    mut connect: F,
293    mut sleep: Sleep,
294) -> Result<S, WorkerError>
295where
296    S: WorkerSession,
297    F: FnMut() -> Fut,
298    Fut: Future<Output = Result<S, WorkerError>>,
299    Sleep: FnMut(Duration) -> SleepFut,
300    SleepFut: Future<Output = ()>,
301{
302    let backoff = ReconnectBackoff::from_config(config)?;
303
304    for attempt in 1..=backoff.attempts() {
305        debug!(attempt, "attempting worker reconnect");
306        let result = match connect().await {
307            Ok(session) => {
308                register_connected_session(
309                    session,
310                    config,
311                    activity_types.clone(),
312                    activities.clone(),
313                    available_handlers,
314                )
315                .await
316            }
317            Err(error) => Err(error),
318        };
319
320        match result {
321            Ok(session) => {
322                debug!(attempt, "worker reconnect succeeded");
323                return Ok(session);
324            }
325            Err(error) => {
326                if !error.is_retryable() {
327                    error!(
328                        attempt,
329                        error = %error,
330                        "worker reconnect denied by server; not retrying"
331                    );
332                    return Err(error);
333                }
334                if attempt == backoff.attempts() {
335                    error!(attempt, error = %error, "worker reconnect attempts exhausted");
336                    return Err(error);
337                }
338                let delay = backoff.delay_for_attempt(attempt);
339                warn!(
340                    attempt,
341                    delay_ms = delay.as_millis(),
342                    error = %error,
343                    "worker reconnect failed; backing off"
344                );
345                sleep(delay).await;
346            }
347        }
348    }
349
350    Err(WorkerError::registration(InvalidReconnectBackoff {
351        message: String::from("reconnect_max_attempts must be greater than zero"),
352    }))
353}
354
355/// Re-reports every unacknowledged result/failure before serving new work.
356///
357/// Server `ResultAck` frames clear entries mid-session, so the steady-state
358/// backlog is empty and this replay decays to the still-unacked residue.
359/// Each send carries the session's per-send deadline.
360///
361/// # Errors
362///
363/// Returns [`WorkerError`] if any re-report send fails. Entries are not removed
364/// by sending; only the explicit `ResultAck` acknowledgement clears the tracker.
365pub async fn re_report_unacked<S>(
366    tracker: &UnackedResultTracker,
367    session: &mut S,
368) -> Result<(), WorkerError>
369where
370    S: WorkerSession,
371{
372    for report in tracker.snapshot() {
373        match report {
374            PendingActivityReport::Completed {
375                workflow_id,
376                activity_id,
377                run_id,
378                completion_token,
379                output,
380            } => {
381                debug!(
382                    workflow_id = %workflow_id,
383                    activity_id = activity_id.sequence_position(),
384                    "re-reporting unacknowledged activity result"
385                );
386                session
387                    .report_result(workflow_id, activity_id, run_id, completion_token, output)
388                    .await?;
389            }
390            PendingActivityReport::Failed {
391                workflow_id,
392                activity_id,
393                run_id,
394                completion_token,
395                failure,
396            } => {
397                debug!(
398                    workflow_id = %workflow_id,
399                    activity_id = activity_id.sequence_position(),
400                    "re-reporting unacknowledged activity failure"
401                );
402                session
403                    .report_failure(workflow_id, activity_id, run_id, completion_token, failure)
404                    .await?;
405            }
406        }
407    }
408    Ok(())
409}
410
411#[derive(Debug, thiserror::Error)]
412#[error("{message}")]
413struct InvalidReconnectBackoff {
414    message: String,
415}
416
417#[cfg(test)]
418mod tests {
419    use std::cell::RefCell;
420    use std::collections::BTreeSet;
421    use std::rc::Rc;
422    use std::time::Duration;
423
424    use aion_core::{
425        ActivityError, ActivityErrorKind, ActivityId, ContentType, Payload, RunId, WorkflowId,
426    };
427    use async_trait::async_trait;
428    use futures::stream;
429
430    use super::{
431        PendingActivityReport, UnackedResultTracker, re_report_unacked, reconnect_with_sleep,
432    };
433    use crate::error::WorkerError;
434    use crate::protocol::{
435        WorkerSession, WorkerSessionEvent, WorkerTaskStream, validate_activity_handlers,
436    };
437    use crate::{ReconnectConfig, WorkerConfig};
438
439    #[test]
440    fn tracker_records_reports_and_acknowledges_by_workflow_and_activity_id() {
441        let workflow_id = WorkflowId::new_v4();
442        let first_id = ActivityId::from_sequence_position(1);
443        let second_id = ActivityId::from_sequence_position(2);
444        let mut tracker = UnackedResultTracker::new();
445
446        tracker.record(PendingActivityReport::Completed {
447            workflow_id: workflow_id.clone(),
448            activity_id: first_id.clone(),
449            run_id: None,
450            completion_token: String::from("generation-1"),
451            output: Payload::new(ContentType::Json, b"{\"first\":true}".to_vec()),
452        });
453        tracker.record(PendingActivityReport::Completed {
454            workflow_id: workflow_id.clone(),
455            activity_id: second_id.clone(),
456            run_id: None,
457            completion_token: String::from("generation-1"),
458            output: Payload::new(ContentType::Json, b"{\"second\":true}".to_vec()),
459        });
460
461        assert_eq!(tracker.len(), 2);
462        assert!(tracker.acknowledge(&workflow_id, &first_id).is_some());
463        assert_eq!(tracker.len(), 1);
464        assert!(tracker.get(&workflow_id, &second_id).is_some());
465        assert!(tracker.get(&workflow_id, &first_id).is_none());
466    }
467
468    #[test]
469    fn tracker_keeps_reports_for_distinct_workflows_at_the_same_sequence_position() {
470        let first_workflow = WorkflowId::new_v4();
471        let second_workflow = WorkflowId::new_v4();
472        let activity_id = ActivityId::from_sequence_position(3);
473        let mut tracker = UnackedResultTracker::new();
474
475        tracker.record(PendingActivityReport::Completed {
476            workflow_id: first_workflow.clone(),
477            activity_id: activity_id.clone(),
478            run_id: None,
479            completion_token: String::from("generation-1"),
480            output: Payload::new(ContentType::Json, b"{\"workflow\":\"a\"}".to_vec()),
481        });
482        tracker.record(PendingActivityReport::Completed {
483            workflow_id: second_workflow.clone(),
484            activity_id: activity_id.clone(),
485            run_id: None,
486            completion_token: String::from("generation-1"),
487            output: Payload::new(ContentType::Json, b"{\"workflow\":\"b\"}".to_vec()),
488        });
489
490        assert_eq!(tracker.len(), 2);
491        assert!(tracker.get(&first_workflow, &activity_id).is_some());
492        assert!(tracker.get(&second_workflow, &activity_id).is_some());
493        assert!(
494            tracker.acknowledge(&first_workflow, &activity_id).is_some(),
495            "acknowledging one workflow's report must not require the other's"
496        );
497        assert_eq!(tracker.len(), 1);
498        assert!(tracker.get(&second_workflow, &activity_id).is_some());
499    }
500
501    #[tokio::test]
502    async fn reconnect_fails_once_then_handshakes_and_registers() -> Result<(), WorkerError> {
503        let config = test_config();
504        let attempts = Rc::new(RefCell::new(0usize));
505        let sleeps = Rc::new(RefCell::new(Vec::new()));
506        let activity_types = vec![String::from("charge-card")];
507        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
508        let attempts_for_connect = Rc::clone(&attempts);
509        let sleeps_for_sleep = Rc::clone(&sleeps);
510
511        let session = reconnect_with_sleep(
512            &config,
513            activity_types.clone(),
514            Vec::new(),
515            &handlers,
516            move || {
517                let attempts_for_connect = Rc::clone(&attempts_for_connect);
518                async move {
519                    let mut attempts = attempts_for_connect.borrow_mut();
520                    *attempts += 1;
521                    if *attempts == 1 {
522                        Err(WorkerError::Transport {
523                            source: tonic::Status::unavailable("disconnected"),
524                        })
525                    } else {
526                        Ok(ReconnectFakeSession::default())
527                    }
528                }
529            },
530            move |delay| {
531                let sleeps_for_sleep = Rc::clone(&sleeps_for_sleep);
532                async move {
533                    sleeps_for_sleep.borrow_mut().push(delay);
534                }
535            },
536        )
537        .await?;
538
539        assert_eq!(*attempts.borrow(), 2);
540        assert_eq!(*sleeps.borrow(), vec![Duration::from_millis(5)]);
541        assert_eq!(session.handshakes, vec![String::from("worker-a")]);
542        assert_eq!(session.registrations, vec![activity_types]);
543        Ok(())
544    }
545
546    #[tokio::test]
547    async fn permission_denied_registration_stops_after_one_attempt() {
548        let config = test_config();
549        let attempts = Rc::new(RefCell::new(0usize));
550        let sleeps = Rc::new(RefCell::new(Vec::new()));
551        let activity_types = vec![String::from("charge-card")];
552        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
553        let attempts_for_connect = Rc::clone(&attempts);
554        let sleeps_for_sleep = Rc::clone(&sleeps);
555
556        let result = reconnect_with_sleep(
557            &config,
558            activity_types,
559            Vec::new(),
560            &handlers,
561            move || {
562                let attempts_for_connect = Rc::clone(&attempts_for_connect);
563                async move {
564                    *attempts_for_connect.borrow_mut() += 1;
565                    Ok(DeniedRegistrationSession {
566                        denial: tonic::Status::permission_denied(
567                            "namespace `payments` is not granted to subject `worker-a`",
568                        ),
569                    })
570                }
571            },
572            move |delay| {
573                let sleeps_for_sleep = Rc::clone(&sleeps_for_sleep);
574                async move {
575                    sleeps_for_sleep.borrow_mut().push(delay);
576                }
577            },
578        )
579        .await;
580
581        assert!(result.is_err());
582        let Err(error) = result else { return };
583        assert_eq!(*attempts.borrow(), 1);
584        assert!(sleeps.borrow().is_empty());
585        assert!(!error.is_retryable());
586        assert!(matches!(
587            error.grpc_status().map(tonic::Status::code),
588            Some(tonic::Code::PermissionDenied)
589        ));
590        assert_eq!(
591            error.grpc_status().map(tonic::Status::message),
592            Some("namespace `payments` is not granted to subject `worker-a`")
593        );
594        assert!(
595            error
596                .to_string()
597                .contains("namespace `payments` is not granted to subject `worker-a`")
598        );
599    }
600
601    #[tokio::test]
602    async fn unauthenticated_handshake_stops_after_one_attempt() {
603        let config = test_config();
604        let attempts = Rc::new(RefCell::new(0usize));
605        let sleeps = Rc::new(RefCell::new(Vec::new()));
606        let activity_types = vec![String::from("charge-card")];
607        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
608        let attempts_for_connect = Rc::clone(&attempts);
609        let sleeps_for_sleep = Rc::clone(&sleeps);
610
611        let result = reconnect_with_sleep(
612            &config,
613            activity_types,
614            Vec::new(),
615            &handlers,
616            move || {
617                let attempts_for_connect = Rc::clone(&attempts_for_connect);
618                async move {
619                    *attempts_for_connect.borrow_mut() += 1;
620                    Err::<ReconnectFakeSession, _>(WorkerError::Handshake {
621                        source: tonic::Status::unauthenticated("worker credentials were rejected"),
622                    })
623                }
624            },
625            move |delay| {
626                let sleeps_for_sleep = Rc::clone(&sleeps_for_sleep);
627                async move {
628                    sleeps_for_sleep.borrow_mut().push(delay);
629                }
630            },
631        )
632        .await;
633
634        assert!(result.is_err());
635        let Err(error) = result else { return };
636        assert_eq!(*attempts.borrow(), 1);
637        assert!(sleeps.borrow().is_empty());
638        assert!(!error.is_retryable());
639        assert!(matches!(
640            error.grpc_status().map(tonic::Status::code),
641            Some(tonic::Code::Unauthenticated)
642        ));
643        assert!(
644            error
645                .to_string()
646                .contains("worker credentials were rejected")
647        );
648    }
649
650    #[tokio::test]
651    async fn unavailable_transport_retries_until_attempts_exhausted() {
652        let config = test_config();
653        let attempts = Rc::new(RefCell::new(0usize));
654        let sleeps = Rc::new(RefCell::new(Vec::new()));
655        let activity_types = vec![String::from("charge-card")];
656        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
657        let attempts_for_connect = Rc::clone(&attempts);
658        let sleeps_for_sleep = Rc::clone(&sleeps);
659
660        let result = reconnect_with_sleep(
661            &config,
662            activity_types,
663            Vec::new(),
664            &handlers,
665            move || {
666                let attempts_for_connect = Rc::clone(&attempts_for_connect);
667                async move {
668                    *attempts_for_connect.borrow_mut() += 1;
669                    Err::<ReconnectFakeSession, _>(WorkerError::Transport {
670                        source: tonic::Status::unavailable("engine unreachable"),
671                    })
672                }
673            },
674            move |delay| {
675                let sleeps_for_sleep = Rc::clone(&sleeps_for_sleep);
676                async move {
677                    sleeps_for_sleep.borrow_mut().push(delay);
678                }
679            },
680        )
681        .await;
682
683        assert!(result.is_err());
684        let Err(error) = result else { return };
685        assert_eq!(*attempts.borrow(), 3);
686        assert_eq!(
687            *sleeps.borrow(),
688            vec![Duration::from_millis(5), Duration::from_millis(10)]
689        );
690        assert!(error.is_retryable());
691        assert!(matches!(
692            error.grpc_status().map(tonic::Status::code),
693            Some(tonic::Code::Unavailable)
694        ));
695    }
696
697    #[tokio::test]
698    async fn re_reports_unacked_reports_without_removing_them() -> Result<(), WorkerError> {
699        let workflow_id = WorkflowId::new_v4();
700        let activity_id = ActivityId::from_sequence_position(7);
701        let output = Payload::new(ContentType::Json, b"{}".to_vec());
702        let mut tracker = UnackedResultTracker::new();
703        tracker.record(PendingActivityReport::Completed {
704            workflow_id: workflow_id.clone(),
705            activity_id: activity_id.clone(),
706            run_id: None,
707            completion_token: String::from("generation-1"),
708            output: output.clone(),
709        });
710        let mut session = ReconnectFakeSession::default();
711
712        re_report_unacked(&tracker, &mut session).await?;
713
714        assert_eq!(tracker.len(), 1);
715        assert_eq!(
716            session.reports,
717            vec![RecordedReport::Completed(workflow_id, activity_id, output)]
718        );
719        Ok(())
720    }
721
722    #[derive(Default)]
723    struct ReconnectFakeSession {
724        handshakes: Vec<String>,
725        registrations: Vec<Vec<String>>,
726        reports: Vec<RecordedReport>,
727    }
728
729    /// Session whose registration is rejected by the server with a gRPC denial,
730    /// mirroring `aion-server` answering `PermissionDenied` for an ungranted
731    /// namespace.
732    struct DeniedRegistrationSession {
733        denial: tonic::Status,
734    }
735
736    #[async_trait]
737    impl WorkerSession for DeniedRegistrationSession {
738        async fn handshake(&mut self, _config: &WorkerConfig) -> Result<(), WorkerError> {
739            Ok(())
740        }
741
742        async fn register(
743            &mut self,
744            activity_types: Vec<String>,
745            available_handlers: &BTreeSet<String>,
746        ) -> Result<(), WorkerError> {
747            validate_activity_handlers(&activity_types, available_handlers)?;
748            Err(WorkerError::Registration {
749                source: Box::new(self.denial.clone()),
750            })
751        }
752
753        fn receive_tasks(&mut self) -> WorkerTaskStream {
754            Box::pin(stream::empty::<Result<WorkerSessionEvent, WorkerError>>())
755        }
756
757        async fn report_result(
758            &mut self,
759            workflow_id: WorkflowId,
760            activity_id: ActivityId,
761            run_id: Option<RunId>,
762            completion_token: String,
763            result: Payload,
764        ) -> Result<(), WorkerError> {
765            drop((workflow_id, activity_id, run_id, completion_token, result));
766            Err(WorkerError::Registration {
767                source: Box::new(self.denial.clone()),
768            })
769        }
770
771        async fn report_failure(
772            &mut self,
773            workflow_id: WorkflowId,
774            activity_id: ActivityId,
775            run_id: Option<RunId>,
776            completion_token: String,
777            failure: ActivityError,
778        ) -> Result<(), WorkerError> {
779            drop((workflow_id, activity_id, run_id, completion_token, failure));
780            Err(WorkerError::Registration {
781                source: Box::new(self.denial.clone()),
782            })
783        }
784
785        async fn send_heartbeat(
786            &mut self,
787            workflow_id: WorkflowId,
788            activity_id: ActivityId,
789            progress: Option<Payload>,
790        ) -> Result<(), WorkerError> {
791            drop((workflow_id, activity_id, progress));
792            Err(WorkerError::Registration {
793                source: Box::new(self.denial.clone()),
794            })
795        }
796    }
797
798    #[derive(Clone, Debug, PartialEq, Eq)]
799    enum RecordedReport {
800        Completed(WorkflowId, ActivityId, Payload),
801        Failed(WorkflowId, ActivityId, ActivityError),
802    }
803
804    #[async_trait]
805    impl WorkerSession for ReconnectFakeSession {
806        async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
807            self.handshakes.push(config.identity.clone());
808            Ok(())
809        }
810
811        async fn register(
812            &mut self,
813            activity_types: Vec<String>,
814            available_handlers: &BTreeSet<String>,
815        ) -> Result<(), WorkerError> {
816            validate_activity_handlers(&activity_types, available_handlers)?;
817            self.registrations.push(activity_types);
818            Ok(())
819        }
820
821        fn receive_tasks(&mut self) -> WorkerTaskStream {
822            Box::pin(stream::empty::<Result<WorkerSessionEvent, WorkerError>>())
823        }
824
825        async fn report_result(
826            &mut self,
827            workflow_id: WorkflowId,
828            activity_id: ActivityId,
829            run_id: Option<RunId>,
830            completion_token: String,
831            result: Payload,
832        ) -> Result<(), WorkerError> {
833            drop((run_id, completion_token));
834            self.reports
835                .push(RecordedReport::Completed(workflow_id, activity_id, result));
836            Ok(())
837        }
838
839        async fn report_failure(
840            &mut self,
841            workflow_id: WorkflowId,
842            activity_id: ActivityId,
843            run_id: Option<RunId>,
844            completion_token: String,
845            failure: ActivityError,
846        ) -> Result<(), WorkerError> {
847            drop((run_id, completion_token));
848            self.reports
849                .push(RecordedReport::Failed(workflow_id, activity_id, failure));
850            Ok(())
851        }
852
853        async fn send_heartbeat(
854            &mut self,
855            workflow_id: WorkflowId,
856            activity_id: ActivityId,
857            progress: Option<Payload>,
858        ) -> Result<(), WorkerError> {
859            drop((workflow_id, activity_id, progress));
860            Ok(())
861        }
862    }
863
864    fn test_config() -> WorkerConfig {
865        WorkerConfig::new(
866            "http://127.0.0.1:50051",
867            "payments",
868            "worker-a",
869            2,
870            ReconnectConfig::new(Duration::from_millis(5), Duration::from_millis(20), 3),
871            None,
872        )
873    }
874
875    fn terminal_failure() -> ActivityError {
876        ActivityError {
877            kind: ActivityErrorKind::Terminal,
878            message: String::from("terminal"),
879            details: None,
880        }
881    }
882
883    #[test]
884    fn tracker_replaces_existing_activity_report() {
885        let workflow_id = WorkflowId::new_v4();
886        let activity_id = ActivityId::from_sequence_position(9);
887        let mut tracker = UnackedResultTracker::new();
888        tracker.record(PendingActivityReport::Completed {
889            workflow_id: workflow_id.clone(),
890            activity_id: activity_id.clone(),
891            run_id: None,
892            completion_token: String::from("generation-1"),
893            output: Payload::new(ContentType::Json, b"{}".to_vec()),
894        });
895        tracker.record(PendingActivityReport::Failed {
896            workflow_id: workflow_id.clone(),
897            activity_id: activity_id.clone(),
898            run_id: None,
899            completion_token: String::from("generation-1"),
900            failure: terminal_failure(),
901        });
902
903        assert_eq!(tracker.len(), 1);
904        assert!(matches!(
905            tracker.get(&workflow_id, &activity_id),
906            Some(PendingActivityReport::Failed { .. })
907        ));
908    }
909}