1use std::collections::BTreeSet;
4use std::future::Future;
5use std::pin::Pin;
6use std::sync::Arc;
7use std::task::Poll;
8
9use serde::Serialize;
10use serde::de::DeserializeOwned;
11use tracing::{error, info, warn};
12
13use crate::activity::{ActivityRegistry, HandlerFuture};
14use crate::config::WorkerConfig;
15use crate::context::ActivityContext;
16use crate::error::WorkerError;
17use crate::protocol::reconnect::{
18 ReconnectBackoff, UnackedResultTracker, re_report_unacked, reconnect_with_backoff,
19 register_connected_session,
20};
21use crate::protocol::{GrpcWorkerSession, WorkerSession};
22use crate::runtime::{
23 NoShutdown, ServeEnd, SessionHealth, serve_activity_tasks, serve_activity_tasks_until,
24};
25
26#[must_use]
28pub struct WorkerBuilder {
29 config: WorkerConfig,
30 activities: ActivityRegistry,
31}
32
33impl WorkerBuilder {
34 pub fn new(config: WorkerConfig) -> Self {
36 Self {
37 config,
38 activities: ActivityRegistry::new(),
39 }
40 }
41
42 pub fn register_activity<Input, Output, Handler>(
48 mut self,
49 activity_type: impl Into<String>,
50 handler: Handler,
51 ) -> Result<Self, WorkerError>
52 where
53 Input: Serialize + DeserializeOwned + Send + Sync + 'static,
54 Output: Serialize + Send + Sync + 'static,
55 Handler: for<'context> Fn(Input, &'context ActivityContext) -> HandlerFuture<'context, Output>
56 + Send
57 + Sync
58 + 'static,
59 {
60 self.activities = self.activities.register_activity(activity_type, handler)?;
61 Ok(self)
62 }
63
64 pub fn register_activity_with_contract<Input, Output, Handler>(
71 mut self,
72 activity_type: impl Into<String>,
73 handler: Handler,
74 ) -> Result<Self, WorkerError>
75 where
76 Input: Serialize + DeserializeOwned + schemars::JsonSchema + Send + Sync + 'static,
77 Output: Serialize + schemars::JsonSchema + Send + Sync + 'static,
78 Handler: for<'context> Fn(Input, &'context ActivityContext) -> HandlerFuture<'context, Output>
79 + Send
80 + Sync
81 + 'static,
82 {
83 self.activities = self
84 .activities
85 .register_activity_with_contract(activity_type, handler)?;
86 Ok(self)
87 }
88
89 pub fn register_activity_with_descriptor<Input, Output, Handler>(
101 mut self,
102 activity_type: impl Into<String>,
103 descriptor: aion_package::ActivityDescriptor,
104 handler: Handler,
105 ) -> Result<Self, WorkerError>
106 where
107 Input: Serialize + DeserializeOwned + Send + Sync + 'static,
108 Output: Serialize + Send + Sync + 'static,
109 Handler: for<'context> Fn(Input, &'context ActivityContext) -> HandlerFuture<'context, Output>
110 + Send
111 + Sync
112 + 'static,
113 {
114 self.activities = self.activities.register_activity_with_descriptor(
115 activity_type,
116 descriptor,
117 handler,
118 )?;
119 Ok(self)
120 }
121
122 pub fn build(self) -> Result<Worker, WorkerError> {
128 if self.activities.is_empty() {
129 return Err(WorkerError::registration(EmptyActivitySet));
130 }
131 let available_handlers = self.activities.activity_types();
132 let activity_types = available_handlers.iter().cloned().collect();
133 let activity_descriptors = self.activities.activity_descriptors();
134 Ok(Worker {
135 config: self.config,
136 activity_types,
137 activity_descriptors,
138 available_handlers,
139 activities: Arc::new(self.activities),
140 })
141 }
142}
143
144#[must_use]
146pub struct Worker {
147 config: WorkerConfig,
148 activity_types: Vec<String>,
149 activity_descriptors: Vec<aion_package::ActivityDescriptor>,
150 available_handlers: BTreeSet<String>,
151 activities: Arc<ActivityRegistry>,
152}
153
154impl Worker {
155 pub fn builder(config: WorkerConfig) -> WorkerBuilder {
157 WorkerBuilder::new(config)
158 }
159
160 #[must_use]
162 pub fn activity_types(&self) -> &[String] {
163 &self.activity_types
164 }
165
166 #[must_use]
168 pub fn activity_descriptors(&self) -> &[aion_package::ActivityDescriptor] {
169 &self.activity_descriptors
170 }
171
172 #[must_use]
174 pub fn available_handlers(&self) -> &BTreeSet<String> {
175 &self.available_handlers
176 }
177
178 fn log_session_established(&self) {
182 info!(
183 identity = %self.config.identity,
184 endpoint = %self.config.endpoint,
185 activity_types = ?self.activity_types,
186 "worker session established; serving activities"
187 );
188 }
189
190 pub async fn run(self) -> Result<(), WorkerError> {
211 self.run_until(std::future::pending::<()>()).await
212 }
213
214 pub async fn run_until<Shutdown>(self, shutdown: Shutdown) -> Result<(), WorkerError>
226 where
227 Shutdown: Future<Output = ()> + Send,
228 {
229 let config = self.config.clone();
230 self.run_with_connector_until(move || GrpcWorkerSession::connect(config.clone()), shutdown)
231 .await
232 }
233
234 pub async fn run_with_connector_until<S, F, Fut, Shutdown>(
275 self,
276 mut connect: F,
277 shutdown: Shutdown,
278 ) -> Result<(), WorkerError>
279 where
280 S: WorkerSession,
281 F: FnMut() -> Fut,
282 Fut: Future<Output = Result<S, WorkerError>>,
283 Shutdown: Future<Output = ()> + Send,
284 {
285 let backoff = ReconnectBackoff::from_config(&self.config)?;
286 let mut tracker = UnackedResultTracker::new();
287 tokio::pin!(shutdown);
288 let mut shutdown = SharedShutdown::new(shutdown);
289 let mut drop_failures = 0_usize;
290 let mut recovery_error: Option<WorkerError> = None;
291
292 loop {
293 let connected = tokio::select! {
294 biased;
295 () = shutdown.wait() => {
296 return recovery_error.take().map_or(Ok(()), Err);
297 }
298 result = reconnect_with_backoff(
299 &self.config,
300 self.activity_types.clone(),
301 self.activity_descriptors.clone(),
302 &self.available_handlers,
303 &mut connect,
304 ) => result,
305 };
306 let mut session = connected?;
307 self.log_session_established();
308 let session_started = tokio::time::Instant::now();
309 let mut health = SessionHealth::default();
310 let replay = tokio::select! {
315 biased;
316 () = shutdown.wait() => None,
317 result = re_report_unacked(&tracker, &mut session) => Some(result),
318 };
319 let Some(replay_result) = replay else {
320 return Ok(());
321 };
322 let served = match replay_result {
323 Ok(()) => {
324 serve_activity_tasks_until(
325 &self.config,
326 &mut session,
327 Arc::clone(&self.activities),
328 &mut tracker,
329 &mut health,
330 shutdown.wait(),
331 )
332 .await
333 }
334 Err(report_error) => Err(report_error),
335 };
336 drop(session);
337 let cause = match classify_serve_outcome(served, &health, shutdown.fired()) {
338 ServeClassification::End(result) => return result,
339 ServeClassification::Drop(cause) => cause,
340 };
341 let connected_for = health
348 .stream_ended_at
349 .unwrap_or_else(tokio::time::Instant::now)
350 .saturating_duration_since(session_started);
351 let proved_healthy = health.tasks_reported > 0 || connected_for > backoff.max_delay();
352 if proved_healthy && drop_failures > 0 {
353 info!(
354 drop_failures,
355 tasks_reported = health.tasks_reported,
356 "worker session proved healthy; drop budget reset"
357 );
358 drop_failures = 0;
359 }
360 let delay = if matches!(cause, DropCause::Drain) {
365 self.config.reconnect.initial_backoff
366 } else {
367 drop_failures += 1;
368 if drop_failures >= backoff.attempts() {
369 let error = cause.into_exhaustion_error();
370 error!(
371 drop_failures,
372 error = %error,
373 "worker session drop budget exhausted; not reconnecting"
374 );
375 return Err(error);
376 }
377 backoff.delay_for_attempt(drop_failures)
378 };
379 warn!(
380 drop_failures,
381 delay_ms = delay.as_millis(),
382 cause = %cause,
383 "worker session dropped; reconnecting after backoff"
384 );
385 let shutdown_won = tokio::select! {
386 biased;
387 () = shutdown.wait() => true,
388 () = tokio::time::sleep(delay) => false,
389 };
390 if shutdown_won {
391 return cause.into_shutdown_result();
392 }
393 recovery_error = cause.into_recovery_error();
394 }
395 }
396
397 pub async fn run_with_session<S>(self, session: S) -> Result<S, WorkerError>
403 where
404 S: WorkerSession,
405 {
406 self.run_with_session_until(session, std::future::pending::<()>())
407 .await
408 }
409
410 pub async fn run_with_session_until<S, Shutdown>(
416 self,
417 session: S,
418 shutdown: Shutdown,
419 ) -> Result<S, WorkerError>
420 where
421 S: WorkerSession,
422 Shutdown: Future<Output = ()> + Send,
423 {
424 let mut session = register_connected_session(
425 session,
426 &self.config,
427 self.activity_types.clone(),
428 self.activity_descriptors.clone(),
429 &self.available_handlers,
430 )
431 .await?;
432 let mut tracker = UnackedResultTracker::new();
433 let mut health = SessionHealth::default();
434 serve_activity_tasks_until(
435 &self.config,
436 &mut session,
437 self.activities,
438 &mut tracker,
439 &mut health,
440 shutdown,
441 )
442 .await?;
443 Ok(session)
444 }
445}
446
447enum ServeClassification {
449 End(Result<(), WorkerError>),
451 Drop(DropCause),
453}
454
455fn classify_serve_outcome(
462 served: Result<ServeEnd, WorkerError>,
463 health: &SessionHealth,
464 shutdown_fired: bool,
465) -> ServeClassification {
466 match served {
467 Ok(ServeEnd::Shutdown) => ServeClassification::End(Ok(())),
468 Ok(ServeEnd::Drained) => {
469 if shutdown_fired {
470 return ServeClassification::End(Ok(()));
471 }
472 ServeClassification::Drop(DropCause::Drain)
473 }
474 Ok(ServeEnd::StreamClosed) => {
475 if shutdown_fired {
476 return ServeClassification::End(Ok(()));
477 }
478 ServeClassification::Drop(DropCause::CleanClose)
479 }
480 Err(error) if !error.is_retryable() => {
481 error!(error = %error, "worker session denied by server; not reconnecting");
482 ServeClassification::End(Err(error))
483 }
484 Err(error) if health.drain_received => {
485 warn!(
489 error = %error,
490 "session error after server drain; classified as drain drop"
491 );
492 if shutdown_fired {
493 return ServeClassification::End(Ok(()));
494 }
495 ServeClassification::Drop(DropCause::Drain)
496 }
497 Err(error) => {
498 if shutdown_fired {
499 return ServeClassification::End(Err(error));
500 }
501 ServeClassification::Drop(DropCause::Failure(error))
502 }
503 }
504}
505
506enum DropCause {
508 Failure(WorkerError),
510 CleanClose,
512 Drain,
515}
516
517impl DropCause {
518 fn into_exhaustion_error(self) -> WorkerError {
524 match self {
525 Self::Failure(error) => error,
526 Self::CleanClose | Self::Drain => WorkerError::CleanCloseExhausted,
527 }
528 }
529
530 fn into_shutdown_result(self) -> Result<(), WorkerError> {
533 match self {
534 Self::Failure(error) => Err(error),
535 Self::CleanClose | Self::Drain => Ok(()),
536 }
537 }
538
539 fn into_recovery_error(self) -> Option<WorkerError> {
541 match self {
542 Self::Failure(error) => Some(error),
543 Self::CleanClose | Self::Drain => None,
544 }
545 }
546}
547
548impl std::fmt::Display for DropCause {
549 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
550 match self {
551 Self::Failure(error) => write!(formatter, "{error}"),
552 Self::CleanClose => write!(formatter, "server closed the worker stream cleanly"),
553 Self::Drain => write!(formatter, "server drained the worker stream"),
554 }
555 }
556}
557
558struct SharedShutdown<'a, S> {
567 inner: Pin<&'a mut S>,
568 fired: bool,
569}
570
571impl<'a, S> SharedShutdown<'a, S>
572where
573 S: Future<Output = ()> + Send,
574{
575 const fn new(inner: Pin<&'a mut S>) -> Self {
576 Self {
577 inner,
578 fired: false,
579 }
580 }
581
582 const fn fired(&self) -> bool {
584 self.fired
585 }
586
587 fn wait(&mut self) -> impl Future<Output = ()> + Send {
589 std::future::poll_fn(|context| {
590 if self.fired {
591 return Poll::Ready(());
592 }
593 match self.inner.as_mut().poll(context) {
594 Poll::Ready(()) => {
595 self.fired = true;
596 Poll::Ready(())
597 }
598 Poll::Pending => Poll::Pending,
599 }
600 })
601 }
602}
603
604pub async fn run_worker_with_session<S>(worker: Worker, session: S) -> Result<S, WorkerError>
610where
611 S: WorkerSession,
612{
613 worker.run_with_session(session).await
614}
615
616#[derive(Debug, thiserror::Error, Clone, PartialEq, Eq)]
618#[error("worker must register at least one activity handler")]
619pub struct EmptyActivitySet;
620
621fn _assert_live_session_type() {
622 let _ = std::mem::size_of::<GrpcWorkerSession>();
623 let _ = std::mem::size_of::<NoShutdown>();
624 let _ = serve_activity_tasks::<GrpcWorkerSession, ActivityRegistry>;
625}
626
627#[cfg(test)]
628mod tests {
629 use std::collections::BTreeSet;
630 use std::sync::Arc;
631 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
632 use std::time::Duration;
633
634 use aion_core::{ActivityError, ActivityId, ContentType, Payload, RunId, WorkflowId};
635 use aion_proto::{ProtoActivityId, ProtoActivityTask, ProtoPayload, ProtoWorkflowId};
636 use async_trait::async_trait;
637 use futures::StreamExt as _;
638 use futures::stream;
639 use serde::{Deserialize, Serialize};
640 use tokio::sync::{Notify, mpsc};
641
642 use super::{Worker, WorkerBuilder};
643 use crate::config::{ReconnectConfig, WorkerConfig};
644 use crate::context::ActivityContext;
645 use crate::error::WorkerError;
646 use crate::protocol::{
647 WorkerSession, WorkerSessionEvent, WorkerTaskStream, validate_activity_handlers,
648 };
649
650 #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
651 struct TestInput {
652 value: i32,
653 }
654
655 #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
656 struct TestOutput {
657 value: i32,
658 }
659
660 struct ChannelSession {
661 receiver: Option<mpsc::Receiver<Result<WorkerSessionEvent, WorkerError>>>,
662 reports: Vec<RecordedReport>,
663 registered: Vec<String>,
664 }
665
666 #[derive(Clone, Debug, PartialEq, Eq)]
667 enum RecordedReport {
668 Completed(WorkflowId, ActivityId, Payload),
669 Failed(WorkflowId, ActivityId, ActivityError),
670 }
671
672 #[async_trait]
673 impl WorkerSession for ChannelSession {
674 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
675 drop(config.clone());
676 Ok(())
677 }
678
679 async fn register(
680 &mut self,
681 activity_types: Vec<String>,
682 available_handlers: &BTreeSet<String>,
683 ) -> Result<(), WorkerError> {
684 validate_activity_handlers(&activity_types, available_handlers)?;
685 self.registered = activity_types;
686 Ok(())
687 }
688
689 fn receive_tasks(&mut self) -> WorkerTaskStream {
690 match self.receiver.take() {
691 Some(receiver) => Box::pin(tokio_stream::wrappers::ReceiverStream::new(receiver)),
692 None => Box::pin(stream::empty()),
693 }
694 }
695
696 async fn report_result(
697 &mut self,
698 workflow_id: WorkflowId,
699 activity_id: ActivityId,
700 run_id: Option<RunId>,
701 completion_token: String,
702 result: Payload,
703 ) -> Result<(), WorkerError> {
704 drop((run_id, completion_token));
705 self.reports
706 .push(RecordedReport::Completed(workflow_id, activity_id, result));
707 Ok(())
708 }
709
710 async fn report_failure(
711 &mut self,
712 workflow_id: WorkflowId,
713 activity_id: ActivityId,
714 run_id: Option<RunId>,
715 completion_token: String,
716 failure: ActivityError,
717 ) -> Result<(), WorkerError> {
718 drop((run_id, completion_token));
719 self.reports
720 .push(RecordedReport::Failed(workflow_id, activity_id, failure));
721 Ok(())
722 }
723
724 async fn send_heartbeat(
725 &mut self,
726 workflow_id: WorkflowId,
727 activity_id: ActivityId,
728 progress: Option<Payload>,
729 ) -> Result<(), WorkerError> {
730 drop((workflow_id, activity_id, progress));
731 Ok(())
732 }
733 }
734
735 struct HungReportSession {
738 log: mpsc::UnboundedSender<SessionLog>,
739 index: usize,
740 }
741
742 #[async_trait]
743 impl WorkerSession for HungReportSession {
744 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
745 drop(config.clone());
746 Ok(())
747 }
748
749 async fn register(
750 &mut self,
751 activity_types: Vec<String>,
752 available_handlers: &BTreeSet<String>,
753 ) -> Result<(), WorkerError> {
754 validate_activity_handlers(&activity_types, available_handlers)?;
755 self.log
756 .send(SessionLog::Registered(self.index, activity_types))
757 .map_err(WorkerError::decode)
758 }
759
760 fn receive_tasks(&mut self) -> WorkerTaskStream {
761 Box::pin(stream::pending())
762 }
763
764 async fn report_result(
765 &mut self,
766 _workflow_id: WorkflowId,
767 _activity_id: ActivityId,
768 _run_id: Option<RunId>,
769 completion_token: String,
770 _result: Payload,
771 ) -> Result<(), WorkerError> {
772 drop(completion_token);
773 std::future::pending::<()>().await;
774 Ok(())
775 }
776
777 async fn report_failure(
778 &mut self,
779 _workflow_id: WorkflowId,
780 _activity_id: ActivityId,
781 _run_id: Option<RunId>,
782 completion_token: String,
783 _failure: ActivityError,
784 ) -> Result<(), WorkerError> {
785 drop(completion_token);
786 std::future::pending::<()>().await;
787 Ok(())
788 }
789
790 async fn send_heartbeat(
791 &mut self,
792 _workflow_id: WorkflowId,
793 _activity_id: ActivityId,
794 _progress: Option<Payload>,
795 ) -> Result<(), WorkerError> {
796 Ok(())
797 }
798 }
799
800 enum SessionKind {
801 Scripted(ScriptedSession),
802 Hung(HungReportSession),
803 }
804
805 #[async_trait]
806 impl WorkerSession for SessionKind {
807 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
808 match self {
809 Self::Scripted(session) => session.handshake(config).await,
810 Self::Hung(session) => session.handshake(config).await,
811 }
812 }
813
814 async fn register(
815 &mut self,
816 activity_types: Vec<String>,
817 available_handlers: &BTreeSet<String>,
818 ) -> Result<(), WorkerError> {
819 match self {
820 Self::Scripted(session) => {
821 session.register(activity_types, available_handlers).await
822 }
823 Self::Hung(session) => session.register(activity_types, available_handlers).await,
824 }
825 }
826
827 fn receive_tasks(&mut self) -> WorkerTaskStream {
828 match self {
829 Self::Scripted(session) => session.receive_tasks(),
830 Self::Hung(session) => session.receive_tasks(),
831 }
832 }
833
834 async fn report_result(
835 &mut self,
836 workflow_id: WorkflowId,
837 activity_id: ActivityId,
838 run_id: Option<RunId>,
839 completion_token: String,
840 result: Payload,
841 ) -> Result<(), WorkerError> {
842 match self {
843 Self::Scripted(session) => {
844 session
845 .report_result(workflow_id, activity_id, run_id, completion_token, result)
846 .await
847 }
848 Self::Hung(session) => {
849 session
850 .report_result(workflow_id, activity_id, run_id, completion_token, result)
851 .await
852 }
853 }
854 }
855
856 async fn report_failure(
857 &mut self,
858 workflow_id: WorkflowId,
859 activity_id: ActivityId,
860 run_id: Option<RunId>,
861 completion_token: String,
862 failure: ActivityError,
863 ) -> Result<(), WorkerError> {
864 match self {
865 Self::Scripted(session) => {
866 session
867 .report_failure(workflow_id, activity_id, run_id, completion_token, failure)
868 .await
869 }
870 Self::Hung(session) => {
871 session
872 .report_failure(workflow_id, activity_id, run_id, completion_token, failure)
873 .await
874 }
875 }
876 }
877
878 async fn send_heartbeat(
879 &mut self,
880 workflow_id: WorkflowId,
881 activity_id: ActivityId,
882 progress: Option<Payload>,
883 ) -> Result<(), WorkerError> {
884 match self {
885 Self::Scripted(session) => {
886 session
887 .send_heartbeat(workflow_id, activity_id, progress)
888 .await
889 }
890 Self::Hung(session) => {
891 session
892 .send_heartbeat(workflow_id, activity_id, progress)
893 .await
894 }
895 }
896 }
897 }
898
899 struct DrainLatchSession {
903 events: Vec<Result<WorkerSessionEvent, WorkerError>>,
904 fail_id: ActivityId,
905 }
906
907 #[async_trait]
908 impl WorkerSession for DrainLatchSession {
909 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
910 drop(config.clone());
911 Ok(())
912 }
913
914 async fn register(
915 &mut self,
916 activity_types: Vec<String>,
917 available_handlers: &BTreeSet<String>,
918 ) -> Result<(), WorkerError> {
919 validate_activity_handlers(&activity_types, available_handlers)
920 }
921
922 fn receive_tasks(&mut self) -> WorkerTaskStream {
923 Box::pin(stream::iter(std::mem::take(&mut self.events)))
924 }
925
926 async fn report_result(
927 &mut self,
928 _workflow_id: WorkflowId,
929 activity_id: ActivityId,
930 _run_id: Option<RunId>,
931 completion_token: String,
932 _result: Payload,
933 ) -> Result<(), WorkerError> {
934 drop(completion_token);
935 if activity_id == self.fail_id {
936 return Err(WorkerError::Transport {
937 source: tonic::Status::unavailable(
938 "stream broke abruptly after the drain frame",
939 ),
940 });
941 }
942 Ok(())
943 }
944
945 async fn report_failure(
946 &mut self,
947 _workflow_id: WorkflowId,
948 _activity_id: ActivityId,
949 _run_id: Option<RunId>,
950 completion_token: String,
951 _failure: ActivityError,
952 ) -> Result<(), WorkerError> {
953 drop(completion_token);
954 Ok(())
955 }
956
957 async fn send_heartbeat(
958 &mut self,
959 _workflow_id: WorkflowId,
960 _activity_id: ActivityId,
961 _progress: Option<Payload>,
962 ) -> Result<(), WorkerError> {
963 Ok(())
964 }
965 }
966
967 enum LatchKind {
968 Latch(DrainLatchSession),
969 Deny(ScriptedSession),
970 }
971
972 #[async_trait]
973 impl WorkerSession for LatchKind {
974 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
975 match self {
976 Self::Latch(session) => session.handshake(config).await,
977 Self::Deny(session) => session.handshake(config).await,
978 }
979 }
980
981 async fn register(
982 &mut self,
983 activity_types: Vec<String>,
984 available_handlers: &BTreeSet<String>,
985 ) -> Result<(), WorkerError> {
986 match self {
987 Self::Latch(session) => session.register(activity_types, available_handlers).await,
988 Self::Deny(session) => session.register(activity_types, available_handlers).await,
989 }
990 }
991
992 fn receive_tasks(&mut self) -> WorkerTaskStream {
993 match self {
994 Self::Latch(session) => session.receive_tasks(),
995 Self::Deny(session) => session.receive_tasks(),
996 }
997 }
998
999 async fn report_result(
1000 &mut self,
1001 workflow_id: WorkflowId,
1002 activity_id: ActivityId,
1003 run_id: Option<RunId>,
1004 completion_token: String,
1005 result: Payload,
1006 ) -> Result<(), WorkerError> {
1007 match self {
1008 Self::Latch(session) => {
1009 session
1010 .report_result(workflow_id, activity_id, run_id, completion_token, result)
1011 .await
1012 }
1013 Self::Deny(session) => {
1014 session
1015 .report_result(workflow_id, activity_id, run_id, completion_token, result)
1016 .await
1017 }
1018 }
1019 }
1020
1021 async fn report_failure(
1022 &mut self,
1023 workflow_id: WorkflowId,
1024 activity_id: ActivityId,
1025 run_id: Option<RunId>,
1026 completion_token: String,
1027 failure: ActivityError,
1028 ) -> Result<(), WorkerError> {
1029 match self {
1030 Self::Latch(session) => {
1031 session
1032 .report_failure(workflow_id, activity_id, run_id, completion_token, failure)
1033 .await
1034 }
1035 Self::Deny(session) => {
1036 session
1037 .report_failure(workflow_id, activity_id, run_id, completion_token, failure)
1038 .await
1039 }
1040 }
1041 }
1042
1043 async fn send_heartbeat(
1044 &mut self,
1045 workflow_id: WorkflowId,
1046 activity_id: ActivityId,
1047 progress: Option<Payload>,
1048 ) -> Result<(), WorkerError> {
1049 match self {
1050 Self::Latch(session) => {
1051 session
1052 .send_heartbeat(workflow_id, activity_id, progress)
1053 .await
1054 }
1055 Self::Deny(session) => {
1056 session
1057 .send_heartbeat(workflow_id, activity_id, progress)
1058 .await
1059 }
1060 }
1061 }
1062 }
1063
1064 #[test]
1065 fn empty_worker_is_rejected() {
1066 let error = WorkerBuilder::new(test_config()).build().err();
1067
1068 assert!(error.is_some_and(|error| error.to_string().contains("at least one activity")));
1069 }
1070
1071 #[test]
1072 fn worker_collects_two_activity_registration_names() -> Result<(), WorkerError> {
1073 let worker = two_activity_worker()?;
1074 let expected = [String::from("double"), String::from("increment")]
1075 .into_iter()
1076 .collect::<BTreeSet<_>>();
1077
1078 assert_eq!(worker.available_handlers(), &expected);
1079 assert_eq!(
1080 worker.activity_types(),
1081 &[String::from("double"), String::from("increment")]
1082 );
1083 Ok(())
1084 }
1085
1086 #[tokio::test]
1087 async fn worker_registers_names_with_session() -> Result<(), WorkerError> {
1088 let worker = two_activity_worker()?;
1089 let session = worker
1090 .run_with_session(ChannelSession {
1091 receiver: None,
1092 reports: Vec::new(),
1093 registered: Vec::new(),
1094 })
1095 .await?;
1096
1097 assert_eq!(
1098 session.registered,
1099 vec![String::from("double"), String::from("increment")]
1100 );
1101 Ok(())
1102 }
1103
1104 #[tokio::test]
1105 async fn shutdown_waits_for_slow_in_flight_activity() -> Result<(), WorkerError> {
1106 let workflow_id = WorkflowId::new_v4();
1107 let activity_id = ActivityId::from_sequence_position(7);
1108 let (sender, receiver) = mpsc::channel(2);
1109 sender
1110 .send(Ok(WorkerSessionEvent::Task(Box::new(proto_task(
1111 workflow_id,
1112 activity_id.clone(),
1113 "slow",
1114 0,
1115 )))))
1116 .await
1117 .map_err(WorkerError::decode)?;
1118 let release = Arc::new(AtomicBool::new(false));
1119 let started = Arc::new(AtomicUsize::new(0));
1120 let worker = Worker::builder(test_config())
1121 .register_activity("slow", {
1122 let release = Arc::clone(&release);
1123 let started = Arc::clone(&started);
1124 move |input: TestInput, context: &ActivityContext| {
1125 let release = Arc::clone(&release);
1126 let started = Arc::clone(&started);
1127 Box::pin(async move {
1128 let _ = input;
1129 started.fetch_add(1, Ordering::SeqCst);
1130 context.cancelled().await;
1131 while !release.load(Ordering::SeqCst) {
1132 tokio::time::sleep(Duration::from_millis(1)).await;
1133 }
1134 Ok(TestOutput { value: 1 })
1135 })
1136 }
1137 })?
1138 .build()?;
1139 let (shutdown_sender, shutdown_receiver) = tokio::sync::oneshot::channel::<()>();
1140 let session = ChannelSession {
1141 receiver: Some(receiver),
1142 reports: Vec::new(),
1143 registered: Vec::new(),
1144 };
1145 let handle = tokio::spawn(async move {
1146 worker
1147 .run_with_session_until(session, async {
1148 let _ = shutdown_receiver.await;
1149 })
1150 .await
1151 });
1152
1153 wait_until_started(&started).await;
1154 shutdown_sender
1155 .send(())
1156 .map_err(|()| WorkerError::decode(SendFailed))?;
1157 tokio::time::sleep(Duration::from_millis(20)).await;
1158 assert!(!handle.is_finished());
1159 release.store(true, Ordering::SeqCst);
1160 drop(sender);
1161 let session = handle.await.map_err(WorkerError::decode)??;
1162
1163 assert_eq!(session.reports.len(), 1);
1164 assert!(matches!(
1165 &session.reports[0],
1166 RecordedReport::Completed(_, reported_id, _) if reported_id == &activity_id
1167 ));
1168 Ok(())
1169 }
1170
1171 fn two_activity_worker() -> Result<Worker, WorkerError> {
1172 two_activity_worker_with(test_config())
1173 }
1174
1175 fn two_activity_worker_with(config: WorkerConfig) -> Result<Worker, WorkerError> {
1176 Worker::builder(config)
1177 .register_activity("double", |input: TestInput, context| {
1178 Box::pin(async move {
1179 let _ = context;
1180 Ok(TestOutput {
1181 value: input.value * 2,
1182 })
1183 })
1184 })?
1185 .register_activity("increment", |input: TestInput, context| {
1186 Box::pin(async move {
1187 let _ = context;
1188 Ok(TestOutput {
1189 value: input.value + 1,
1190 })
1191 })
1192 })?
1193 .build()
1194 }
1195
1196 fn proto_task(
1197 workflow_id: WorkflowId,
1198 activity_id: ActivityId,
1199 activity_type: &str,
1200 value: i32,
1201 ) -> ProtoActivityTask {
1202 ProtoActivityTask {
1203 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
1204 activity_id: Some(ProtoActivityId::from(activity_id)),
1205 run_id: Some(aion_proto::ProtoRunId::from(RunId::new_v4())),
1206 activity_type: activity_type.to_owned(),
1207 input: Some(ProtoPayload::from(Payload::new(
1208 ContentType::Json,
1209 format!("{{\"value\":{value}}}").into_bytes(),
1210 ))),
1211 attempt: 1,
1212 labels: std::collections::HashMap::new(),
1213 completion_token: String::from("generation-1"),
1214 idempotency_key: String::from("effect-key"),
1215 }
1216 }
1217
1218 async fn wait_until_started(started: &AtomicUsize) {
1219 while started.load(Ordering::SeqCst) == 0 {
1220 tokio::time::sleep(Duration::from_millis(1)).await;
1221 }
1222 }
1223
1224 fn test_config() -> WorkerConfig {
1225 test_config_with(ReconnectConfig::new(
1226 Duration::from_millis(5),
1227 Duration::from_millis(20),
1228 3,
1229 ))
1230 }
1231
1232 fn test_config_with(reconnect: ReconnectConfig) -> WorkerConfig {
1233 WorkerConfig::new(
1234 "http://127.0.0.1:50051",
1235 "payments",
1236 "worker-a",
1237 1,
1238 reconnect,
1239 None,
1240 )
1241 }
1242
1243 fn slow_reconnect_config() -> WorkerConfig {
1244 test_config_with(ReconnectConfig::new(
1245 Duration::from_secs(5),
1246 Duration::from_secs(10),
1247 5,
1248 ))
1249 }
1250
1251 #[derive(Debug, thiserror::Error)]
1252 #[error("failed to send shutdown signal")]
1253 struct SendFailed;
1254
1255 #[derive(Debug, thiserror::Error)]
1256 #[error("expected the worker run to fail")]
1257 struct UnexpectedSuccess;
1258
1259 #[derive(Debug, thiserror::Error)]
1260 #[error("expected a completed activity report")]
1261 struct UnexpectedReportShape;
1262
1263 #[derive(Debug)]
1265 enum SessionLog {
1266 Registered(usize, Vec<String>),
1267 Reported(usize, RecordedReport),
1268 }
1269
1270 struct ScriptedSession {
1273 index: usize,
1274 log: mpsc::UnboundedSender<SessionLog>,
1275 events: Vec<Result<WorkerSessionEvent, WorkerError>>,
1276 fail_reports: bool,
1277 register_denial: Option<tonic::Status>,
1278 delay_stream: Option<Duration>,
1281 }
1282
1283 #[async_trait]
1284 impl WorkerSession for ScriptedSession {
1285 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
1286 drop(config.clone());
1287 Ok(())
1288 }
1289
1290 async fn register(
1291 &mut self,
1292 activity_types: Vec<String>,
1293 available_handlers: &BTreeSet<String>,
1294 ) -> Result<(), WorkerError> {
1295 validate_activity_handlers(&activity_types, available_handlers)?;
1296 if let Some(denial) = self.register_denial.take() {
1297 return Err(WorkerError::Registration {
1298 source: Box::new(denial),
1299 });
1300 }
1301 self.log
1302 .send(SessionLog::Registered(self.index, activity_types))
1303 .map_err(WorkerError::decode)
1304 }
1305
1306 fn receive_tasks(&mut self) -> WorkerTaskStream {
1307 let events = std::mem::take(&mut self.events);
1308 match self.delay_stream.take() {
1309 Some(delay) => Box::pin(
1310 stream::once(async move {
1311 tokio::time::sleep(delay).await;
1312 stream::iter(events)
1313 })
1314 .flatten(),
1315 ),
1316 None => Box::pin(stream::iter(events)),
1317 }
1318 }
1319
1320 async fn report_result(
1321 &mut self,
1322 workflow_id: WorkflowId,
1323 activity_id: ActivityId,
1324 run_id: Option<RunId>,
1325 completion_token: String,
1326 result: Payload,
1327 ) -> Result<(), WorkerError> {
1328 drop((run_id, completion_token));
1329 if self.fail_reports {
1330 return Err(WorkerError::Transport {
1331 source: tonic::Status::unavailable("transport dropped before result ack"),
1332 });
1333 }
1334 self.log
1335 .send(SessionLog::Reported(
1336 self.index,
1337 RecordedReport::Completed(workflow_id, activity_id, result),
1338 ))
1339 .map_err(WorkerError::decode)
1340 }
1341
1342 async fn report_failure(
1343 &mut self,
1344 workflow_id: WorkflowId,
1345 activity_id: ActivityId,
1346 run_id: Option<RunId>,
1347 completion_token: String,
1348 failure: ActivityError,
1349 ) -> Result<(), WorkerError> {
1350 drop((run_id, completion_token));
1351 if self.fail_reports {
1352 return Err(WorkerError::Transport {
1353 source: tonic::Status::unavailable("transport dropped before failure ack"),
1354 });
1355 }
1356 self.log
1357 .send(SessionLog::Reported(
1358 self.index,
1359 RecordedReport::Failed(workflow_id, activity_id, failure),
1360 ))
1361 .map_err(WorkerError::decode)
1362 }
1363
1364 async fn send_heartbeat(
1365 &mut self,
1366 workflow_id: WorkflowId,
1367 activity_id: ActivityId,
1368 progress: Option<Payload>,
1369 ) -> Result<(), WorkerError> {
1370 drop((workflow_id, activity_id, progress));
1371 Ok(())
1372 }
1373 }
1374
1375 #[tokio::test]
1376 async fn establishment_retries_transient_failures_until_attempts_exhausted()
1377 -> Result<(), WorkerError> {
1378 let worker = two_activity_worker()?;
1379 let attempts = Arc::new(AtomicUsize::new(0));
1380 let connect = {
1381 let attempts = Arc::clone(&attempts);
1382 move || {
1383 attempts.fetch_add(1, Ordering::SeqCst);
1384 async move {
1385 Err::<ScriptedSession, _>(WorkerError::Transport {
1386 source: tonic::Status::unavailable("engine unreachable"),
1387 })
1388 }
1389 }
1390 };
1391
1392 let result = worker
1393 .run_with_connector_until(connect, std::future::pending::<()>())
1394 .await;
1395
1396 assert_eq!(attempts.load(Ordering::SeqCst), 3);
1397 let Err(error) = result else {
1398 return Err(WorkerError::decode(UnexpectedSuccess));
1399 };
1400 assert!(error.is_retryable());
1401 assert!(matches!(
1402 error.grpc_status().map(tonic::Status::code),
1403 Some(tonic::Code::Unavailable)
1404 ));
1405 Ok(())
1406 }
1407
1408 #[tokio::test]
1409 async fn establishment_denial_surfaces_after_one_attempt() -> Result<(), WorkerError> {
1410 let worker = two_activity_worker()?;
1411 let attempts = Arc::new(AtomicUsize::new(0));
1412 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1413 let connect = {
1414 let attempts = Arc::clone(&attempts);
1415 move || {
1416 attempts.fetch_add(1, Ordering::SeqCst);
1417 let log = log_sender.clone();
1418 async move {
1419 Ok(ScriptedSession {
1420 index: 1,
1421 log,
1422 events: Vec::new(),
1423 fail_reports: false,
1424 register_denial: Some(tonic::Status::permission_denied(
1425 "namespace `payments` is not granted to subject `worker-a`",
1426 )),
1427 delay_stream: None,
1428 })
1429 }
1430 }
1431 };
1432
1433 let result = worker
1434 .run_with_connector_until(connect, std::future::pending::<()>())
1435 .await;
1436
1437 assert_eq!(attempts.load(Ordering::SeqCst), 1);
1438 let Err(error) = result else {
1439 return Err(WorkerError::decode(UnexpectedSuccess));
1440 };
1441 assert!(!error.is_retryable());
1442 assert!(matches!(
1443 error.grpc_status().map(tonic::Status::code),
1444 Some(tonic::Code::PermissionDenied)
1445 ));
1446 assert_eq!(
1447 error.grpc_status().map(tonic::Status::message),
1448 Some("namespace `payments` is not granted to subject `worker-a`")
1449 );
1450 drop(log_receiver);
1451 Ok(())
1452 }
1453
1454 #[tokio::test]
1455 async fn mid_run_drop_reconnects_re_registers_and_re_reports_unacked() -> Result<(), WorkerError>
1456 {
1457 let workflow_id = WorkflowId::new_v4();
1458 let activity_id = ActivityId::from_sequence_position(3);
1459 let worker = two_activity_worker()?;
1460 let attempts = Arc::new(AtomicUsize::new(0));
1461 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
1462 let connect = {
1463 let attempts = Arc::clone(&attempts);
1464 let log_sender = log_sender.clone();
1465 let workflow_id = workflow_id.clone();
1466 let activity_id = activity_id.clone();
1467 move || {
1468 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1469 let log = log_sender.clone();
1470 let task = proto_task(workflow_id.clone(), activity_id.clone(), "double", 21);
1471 async move {
1472 if attempt == 1 {
1473 Ok(ScriptedSession {
1474 index: 1,
1475 log,
1476 events: vec![Ok(WorkerSessionEvent::Task(Box::new(task)))],
1477 fail_reports: true,
1478 register_denial: None,
1479 delay_stream: None,
1480 })
1481 } else if attempt == 2 {
1482 Ok(ScriptedSession {
1483 index: attempt,
1484 log,
1485 events: Vec::new(),
1486 fail_reports: false,
1487 register_denial: None,
1488 delay_stream: None,
1489 })
1490 } else {
1491 Ok(ScriptedSession {
1494 index: attempt,
1495 log,
1496 events: Vec::new(),
1497 fail_reports: false,
1498 register_denial: Some(tonic::Status::permission_denied(
1499 "namespace `payments` revoked for subject `worker-a`",
1500 )),
1501 delay_stream: None,
1502 })
1503 }
1504 }
1505 }
1506 };
1507
1508 let result = worker
1509 .run_with_connector_until(connect, std::future::pending::<()>())
1510 .await;
1511
1512 drop(log_sender);
1513 let mut registrations = Vec::new();
1514 let mut reports = Vec::new();
1515 while let Some(entry) = log_receiver.recv().await {
1516 match entry {
1517 SessionLog::Registered(index, types) => registrations.push((index, types)),
1518 SessionLog::Reported(index, report) => reports.push((index, report)),
1519 }
1520 }
1521 let Err(error) = result else {
1522 return Err(WorkerError::decode(UnexpectedSuccess));
1523 };
1524 assert!(!error.is_retryable());
1525 assert_eq!(attempts.load(Ordering::SeqCst), 3);
1526 let expected_types = vec![String::from("double"), String::from("increment")];
1527 assert_eq!(
1528 registrations,
1529 vec![(1, expected_types.clone()), (2, expected_types)]
1530 );
1531 assert_eq!(reports.len(), 1);
1532 let (session_index, report) = &reports[0];
1533 assert_eq!(*session_index, 2);
1534 let RecordedReport::Completed(reported_workflow, reported_id, payload) = report else {
1535 return Err(WorkerError::decode(UnexpectedReportShape));
1536 };
1537 assert_eq!(reported_workflow, &workflow_id);
1538 assert_eq!(reported_id, &activity_id);
1539 let output: TestOutput =
1540 serde_json::from_slice(payload.bytes()).map_err(WorkerError::decode)?;
1541 assert_eq!(output.value, 42);
1542 Ok(())
1543 }
1544
1545 #[tokio::test]
1546 async fn mid_run_drop_re_reports_unacked_results_for_all_workflows() -> Result<(), WorkerError>
1547 {
1548 let first_workflow = WorkflowId::new_v4();
1549 let second_workflow = WorkflowId::new_v4();
1550 let activity_id = ActivityId::from_sequence_position(3);
1551 let worker = two_activity_worker()?;
1552 let attempts = Arc::new(AtomicUsize::new(0));
1553 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
1554 let connect = {
1555 let attempts = Arc::clone(&attempts);
1556 let log_sender = log_sender.clone();
1557 let first_workflow = first_workflow.clone();
1558 let second_workflow = second_workflow.clone();
1559 let activity_id = activity_id.clone();
1560 move || {
1561 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1562 let log = log_sender.clone();
1563 let first_task =
1564 proto_task(first_workflow.clone(), activity_id.clone(), "double", 10);
1565 let second_task =
1566 proto_task(second_workflow.clone(), activity_id.clone(), "double", 20);
1567 async move {
1568 if attempt == 1 {
1569 Ok(ScriptedSession {
1570 index: 1,
1571 log,
1572 events: vec![
1573 Ok(WorkerSessionEvent::Task(Box::new(first_task))),
1574 Ok(WorkerSessionEvent::Task(Box::new(second_task))),
1575 ],
1576 fail_reports: true,
1577 register_denial: None,
1578 delay_stream: None,
1579 })
1580 } else if attempt == 2 {
1581 Ok(ScriptedSession {
1582 index: attempt,
1583 log,
1584 events: Vec::new(),
1585 fail_reports: false,
1586 register_denial: None,
1587 delay_stream: None,
1588 })
1589 } else {
1590 Ok(ScriptedSession {
1593 index: attempt,
1594 log,
1595 events: Vec::new(),
1596 fail_reports: false,
1597 register_denial: Some(tonic::Status::permission_denied(
1598 "namespace `payments` revoked for subject `worker-a`",
1599 )),
1600 delay_stream: None,
1601 })
1602 }
1603 }
1604 }
1605 };
1606
1607 let result = worker
1608 .run_with_connector_until(connect, std::future::pending::<()>())
1609 .await;
1610
1611 drop(log_sender);
1612 let mut reports = Vec::new();
1613 while let Some(entry) = log_receiver.recv().await {
1614 if let SessionLog::Reported(index, report) = entry {
1615 reports.push((index, report));
1616 }
1617 }
1618 let Err(error) = result else {
1619 return Err(WorkerError::decode(UnexpectedSuccess));
1620 };
1621 assert!(!error.is_retryable());
1622 assert_eq!(attempts.load(Ordering::SeqCst), 3);
1623 assert_eq!(
1624 reports.len(),
1625 2,
1626 "both workflows' colliding sequence-position results must be re-reported"
1627 );
1628 let mut reported_workflows = Vec::new();
1629 for (session_index, report) in &reports {
1630 assert_eq!(*session_index, 2, "re-reports must land on the new session");
1631 let RecordedReport::Completed(reported_workflow, reported_id, _) = report else {
1632 return Err(WorkerError::decode(UnexpectedReportShape));
1633 };
1634 assert_eq!(reported_id, &activity_id);
1635 reported_workflows.push(reported_workflow.clone());
1636 }
1637 assert!(reported_workflows.contains(&first_workflow));
1638 assert!(reported_workflows.contains(&second_workflow));
1639 Ok(())
1640 }
1641
1642 #[tokio::test]
1643 async fn shutdown_during_recovery_establishment_returns_original_drop_error()
1644 -> Result<(), WorkerError> {
1645 let worker = two_activity_worker()?;
1646 let attempts = Arc::new(AtomicUsize::new(0));
1647 let notify = Arc::new(Notify::new());
1648 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1649 let connect = {
1650 let attempts = Arc::clone(&attempts);
1651 let notify = Arc::clone(¬ify);
1652 move || {
1653 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1654 let notify = Arc::clone(¬ify);
1655 let log = log_sender.clone();
1656 async move {
1657 if attempt == 1 {
1658 Ok(ScriptedSession {
1659 index: 1,
1660 log,
1661 events: vec![Err(WorkerError::Transport {
1662 source: tonic::Status::unavailable("stream reset by peer"),
1663 })],
1664 fail_reports: false,
1665 register_denial: None,
1666 delay_stream: None,
1667 })
1668 } else {
1669 notify.notify_one();
1673 std::future::pending::<()>().await;
1674 Err(WorkerError::Transport {
1675 source: tonic::Status::unavailable("unreachable"),
1676 })
1677 }
1678 }
1679 }
1680 };
1681 let shutdown = {
1682 let notify = Arc::clone(¬ify);
1683 async move {
1684 notify.notified().await;
1685 }
1686 };
1687
1688 let run = worker.run_with_connector_until(connect, shutdown);
1689 let result = tokio::time::timeout(Duration::from_secs(5), run)
1690 .await
1691 .map_err(WorkerError::decode)?;
1692
1693 assert_eq!(attempts.load(Ordering::SeqCst), 2);
1694 let Err(error) = result else {
1695 return Err(WorkerError::decode(UnexpectedSuccess));
1696 };
1697 assert!(matches!(
1698 error.grpc_status().map(tonic::Status::code),
1699 Some(tonic::Code::Unavailable)
1700 ));
1701 assert_eq!(
1702 error.grpc_status().map(tonic::Status::message),
1703 Some("stream reset by peer"),
1704 "shutdown during recovery establishment must surface the original drop error"
1705 );
1706 drop(log_receiver);
1707 Ok(())
1708 }
1709
1710 #[tokio::test(start_paused = true)]
1714 async fn mid_run_drop_budget_exhaustion_surfaces_last_drop_error() -> Result<(), WorkerError> {
1715 let worker = two_activity_worker()?;
1716 let attempts = Arc::new(AtomicUsize::new(0));
1717 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1718 let connect = {
1719 let attempts = Arc::clone(&attempts);
1720 move || {
1721 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1722 let log = log_sender.clone();
1723 async move {
1724 Ok(ScriptedSession {
1725 index: attempt,
1726 log,
1727 events: vec![Err(WorkerError::Transport {
1728 source: tonic::Status::unavailable("stream reset by peer"),
1729 })],
1730 fail_reports: false,
1731 register_denial: None,
1732 delay_stream: None,
1733 })
1734 }
1735 }
1736 };
1737
1738 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
1739 let result = tokio::time::timeout(Duration::from_secs(5), run)
1740 .await
1741 .map_err(WorkerError::decode)?;
1742
1743 assert_eq!(attempts.load(Ordering::SeqCst), 3);
1746 let Err(error) = result else {
1747 return Err(WorkerError::decode(UnexpectedSuccess));
1748 };
1749 assert!(error.is_retryable());
1750 assert!(matches!(
1751 error.grpc_status().map(tonic::Status::code),
1752 Some(tonic::Code::Unavailable)
1753 ));
1754 assert_eq!(
1755 error.grpc_status().map(tonic::Status::message),
1756 Some("stream reset by peer")
1757 );
1758 drop(log_receiver);
1759 Ok(())
1760 }
1761
1762 #[tokio::test]
1763 async fn mid_run_denial_surfaces_without_reconnect() -> Result<(), WorkerError> {
1764 let worker = two_activity_worker()?;
1765 let attempts = Arc::new(AtomicUsize::new(0));
1766 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1767 let connect = {
1768 let attempts = Arc::clone(&attempts);
1769 move || {
1770 attempts.fetch_add(1, Ordering::SeqCst);
1771 let log = log_sender.clone();
1772 async move {
1773 Ok(ScriptedSession {
1774 index: 1,
1775 log,
1776 events: vec![Err(WorkerError::Transport {
1777 source: tonic::Status::permission_denied(
1778 "namespace `payments` revoked for subject `worker-a`",
1779 ),
1780 })],
1781 fail_reports: false,
1782 register_denial: None,
1783 delay_stream: None,
1784 })
1785 }
1786 }
1787 };
1788
1789 let result = worker
1790 .run_with_connector_until(connect, std::future::pending::<()>())
1791 .await;
1792
1793 assert_eq!(attempts.load(Ordering::SeqCst), 1);
1794 let Err(error) = result else {
1795 return Err(WorkerError::decode(UnexpectedSuccess));
1796 };
1797 assert!(!error.is_retryable());
1798 assert!(matches!(
1799 error.grpc_status().map(tonic::Status::code),
1800 Some(tonic::Code::PermissionDenied)
1801 ));
1802 assert_eq!(
1803 error.grpc_status().map(tonic::Status::message),
1804 Some("namespace `payments` revoked for subject `worker-a`")
1805 );
1806 drop(log_receiver);
1807 Ok(())
1808 }
1809
1810 #[tokio::test]
1811 async fn shutdown_during_establishment_backoff_returns_promptly() -> Result<(), WorkerError> {
1812 let worker = two_activity_worker_with(slow_reconnect_config())?;
1813 let attempts = Arc::new(AtomicUsize::new(0));
1814 let notify = Arc::new(Notify::new());
1815 let connect = {
1816 let attempts = Arc::clone(&attempts);
1817 let notify = Arc::clone(¬ify);
1818 move || {
1819 attempts.fetch_add(1, Ordering::SeqCst);
1820 notify.notify_one();
1821 async move {
1822 Err::<ScriptedSession, _>(WorkerError::Transport {
1823 source: tonic::Status::unavailable("engine unreachable"),
1824 })
1825 }
1826 }
1827 };
1828 let shutdown = {
1829 let notify = Arc::clone(¬ify);
1830 async move {
1831 notify.notified().await;
1832 }
1833 };
1834
1835 let run = worker.run_with_connector_until(connect, shutdown);
1836 tokio::time::timeout(Duration::from_millis(500), run)
1837 .await
1838 .map_err(WorkerError::decode)??;
1839
1840 assert_eq!(attempts.load(Ordering::SeqCst), 1);
1841 Ok(())
1842 }
1843
1844 #[tokio::test]
1845 async fn shutdown_during_mid_run_drop_backoff_returns_promptly() -> Result<(), WorkerError> {
1846 let worker = two_activity_worker_with(slow_reconnect_config())?;
1847 let attempts = Arc::new(AtomicUsize::new(0));
1848 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1849 let connect = {
1850 let attempts = Arc::clone(&attempts);
1851 move || {
1852 attempts.fetch_add(1, Ordering::SeqCst);
1853 let log = log_sender.clone();
1854 async move {
1855 Ok(ScriptedSession {
1856 index: 1,
1857 log,
1858 events: vec![Err(WorkerError::Transport {
1859 source: tonic::Status::unavailable("stream reset by peer"),
1860 })],
1861 fail_reports: false,
1862 register_denial: None,
1863 delay_stream: None,
1864 })
1865 }
1866 }
1867 };
1868 let shutdown = async {
1869 tokio::time::sleep(Duration::from_millis(50)).await;
1870 };
1871
1872 let run = worker.run_with_connector_until(connect, shutdown);
1873 let result = tokio::time::timeout(Duration::from_millis(500), run)
1874 .await
1875 .map_err(WorkerError::decode)?;
1876
1877 assert_eq!(attempts.load(Ordering::SeqCst), 1);
1878 let Err(error) = result else {
1879 return Err(WorkerError::decode(UnexpectedSuccess));
1880 };
1881 assert!(error.is_retryable());
1882 assert!(matches!(
1883 error.grpc_status().map(tonic::Status::code),
1884 Some(tonic::Code::Unavailable)
1885 ));
1886 drop(log_receiver);
1887 Ok(())
1888 }
1889
1890 #[tokio::test]
1891 async fn served_tasks_reset_drop_budget_across_cycles() -> Result<(), WorkerError> {
1892 let workflow_id = WorkflowId::new_v4();
1893 let activity_id = ActivityId::from_sequence_position(7);
1894 let worker = two_activity_worker_with(test_config_with(ReconnectConfig::new(
1897 Duration::from_millis(1),
1898 Duration::from_secs(3600),
1899 2,
1900 )))?;
1901 let attempts = Arc::new(AtomicUsize::new(0));
1902 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
1903 let connect = {
1904 let attempts = Arc::clone(&attempts);
1905 let log_sender = log_sender.clone();
1906 let workflow_id = workflow_id.clone();
1907 let activity_id = activity_id.clone();
1908 move || {
1909 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1910 let log = log_sender.clone();
1911 let task = proto_task(workflow_id.clone(), activity_id.clone(), "double", 21);
1912 async move {
1913 if attempt <= 4 {
1914 Ok(ScriptedSession {
1915 index: attempt,
1916 log,
1917 events: vec![
1918 Ok(WorkerSessionEvent::Task(Box::new(task))),
1919 Err(WorkerError::Transport {
1920 source: tonic::Status::unavailable("stream reset by peer"),
1921 }),
1922 ],
1923 fail_reports: false,
1924 register_denial: None,
1925 delay_stream: None,
1926 })
1927 } else {
1928 Ok(ScriptedSession {
1929 index: attempt,
1930 log,
1931 events: Vec::new(),
1932 fail_reports: false,
1933 register_denial: Some(tonic::Status::permission_denied(
1934 "namespace `payments` revoked for subject `worker-a`",
1935 )),
1936 delay_stream: None,
1937 })
1938 }
1939 }
1940 }
1941 };
1942
1943 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
1944 let result = tokio::time::timeout(Duration::from_secs(5), run)
1945 .await
1946 .map_err(WorkerError::decode)?;
1947
1948 drop(log_sender);
1949 let mut registrations = 0_usize;
1950 while let Some(entry) = log_receiver.recv().await {
1951 if let SessionLog::Registered(..) = entry {
1952 registrations += 1;
1953 }
1954 }
1955 assert_eq!(attempts.load(Ordering::SeqCst), 5);
1960 assert_eq!(registrations, 4);
1961 let Err(error) = result else {
1962 return Err(WorkerError::decode(UnexpectedSuccess));
1963 };
1964 assert!(!error.is_retryable());
1965 assert!(matches!(
1966 error.grpc_status().map(tonic::Status::code),
1967 Some(tonic::Code::PermissionDenied)
1968 ));
1969 Ok(())
1970 }
1971
1972 #[tokio::test(start_paused = true)]
1973 async fn session_outliving_max_backoff_resets_drop_budget() -> Result<(), WorkerError> {
1974 let worker = two_activity_worker_with(test_config_with(ReconnectConfig::new(
1975 Duration::from_millis(5),
1976 Duration::from_millis(20),
1977 2,
1978 )))?;
1979 let attempts = Arc::new(AtomicUsize::new(0));
1980 let (log_sender, log_receiver) = mpsc::unbounded_channel();
1981 let connect = {
1982 let attempts = Arc::clone(&attempts);
1983 move || {
1984 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
1985 let log = log_sender.clone();
1986 async move {
1987 Ok(ScriptedSession {
1988 index: attempt,
1989 log,
1990 events: vec![Err(WorkerError::Transport {
1991 source: tonic::Status::unavailable("stream reset by peer"),
1992 })],
1993 fail_reports: false,
1994 register_denial: None,
1995 delay_stream: (attempt == 2).then_some(Duration::from_millis(30)),
1999 })
2000 }
2001 }
2002 };
2003
2004 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
2005 let result = tokio::time::timeout(Duration::from_secs(5), run)
2006 .await
2007 .map_err(WorkerError::decode)?;
2008
2009 assert_eq!(attempts.load(Ordering::SeqCst), 3);
2016 let Err(error) = result else {
2017 return Err(WorkerError::decode(UnexpectedSuccess));
2018 };
2019 assert!(error.is_retryable());
2020 assert!(matches!(
2021 error.grpc_status().map(tonic::Status::code),
2022 Some(tonic::Code::Unavailable)
2023 ));
2024 drop(log_receiver);
2025 Ok(())
2026 }
2027
2028 #[tokio::test(start_paused = true)]
2035 async fn post_drop_drain_time_does_not_reset_drop_budget() -> Result<(), WorkerError> {
2036 let workflow_id = WorkflowId::new_v4();
2037 let activity_id = ActivityId::from_sequence_position(9);
2038 let config = WorkerConfig::new(
2041 "http://127.0.0.1:50051",
2042 "payments",
2043 "worker-a",
2044 2,
2045 ReconnectConfig::new(Duration::from_millis(5), Duration::from_millis(20), 2),
2046 None,
2047 );
2048 let worker = Worker::builder(config)
2049 .register_activity("slow", |input: TestInput, context: &ActivityContext| {
2050 let _ = (input, context);
2051 Box::pin(async move {
2052 tokio::time::sleep(Duration::from_millis(60)).await;
2055 Ok(TestOutput { value: 1 })
2056 })
2057 })?
2058 .build()?;
2059 let attempts = Arc::new(AtomicUsize::new(0));
2060 let (log_sender, log_receiver) = mpsc::unbounded_channel();
2061 let connect = {
2062 let attempts = Arc::clone(&attempts);
2063 let workflow_id = workflow_id.clone();
2064 let activity_id = activity_id.clone();
2065 move || {
2066 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
2067 let log = log_sender.clone();
2068 let task = proto_task(workflow_id.clone(), activity_id.clone(), "slow", 1);
2069 async move {
2070 if attempt == 1 {
2071 Ok(ScriptedSession {
2075 index: 1,
2076 log,
2077 events: vec![Err(WorkerError::Transport {
2078 source: tonic::Status::unavailable("stream reset by peer"),
2079 })],
2080 fail_reports: false,
2081 register_denial: None,
2082 delay_stream: None,
2083 })
2084 } else {
2085 Ok(ScriptedSession {
2090 index: attempt,
2091 log,
2092 events: vec![
2093 Ok(WorkerSessionEvent::Task(Box::new(task))),
2094 Err(WorkerError::Transport {
2095 source: tonic::Status::unavailable("stream reset by peer"),
2096 }),
2097 ],
2098 fail_reports: true,
2099 register_denial: None,
2100 delay_stream: None,
2101 })
2102 }
2103 }
2104 }
2105 };
2106
2107 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
2108 let result = tokio::time::timeout(Duration::from_secs(5), run)
2109 .await
2110 .map_err(WorkerError::decode)?;
2111
2112 assert_eq!(attempts.load(Ordering::SeqCst), 2);
2119 let Err(error) = result else {
2120 return Err(WorkerError::decode(UnexpectedSuccess));
2121 };
2122 assert!(error.is_retryable());
2123 assert!(matches!(
2124 error.grpc_status().map(tonic::Status::code),
2125 Some(tonic::Code::Unavailable)
2126 ));
2127 drop(log_receiver);
2128 Ok(())
2129 }
2130
2131 #[tokio::test]
2132 async fn clean_close_reconnects_re_registers_and_keeps_serving() -> Result<(), WorkerError> {
2133 let workflow_id = WorkflowId::new_v4();
2134 let first_activity = ActivityId::from_sequence_position(1);
2135 let second_activity = ActivityId::from_sequence_position(2);
2136 let worker = two_activity_worker()?;
2137 let attempts = Arc::new(AtomicUsize::new(0));
2138 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
2139 let connect = {
2140 let attempts = Arc::clone(&attempts);
2141 let log_sender = log_sender.clone();
2142 let workflow_id = workflow_id.clone();
2143 let first_activity = first_activity.clone();
2144 let second_activity = second_activity.clone();
2145 move || {
2146 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
2147 let log = log_sender.clone();
2148 let first_task =
2149 proto_task(workflow_id.clone(), first_activity.clone(), "double", 10);
2150 let second_task =
2151 proto_task(workflow_id.clone(), second_activity.clone(), "double", 20);
2152 async move {
2153 match attempt {
2154 1 => Ok(ScriptedSession {
2157 index: 1,
2158 log,
2159 events: vec![Ok(WorkerSessionEvent::Task(Box::new(first_task)))],
2160 fail_reports: false,
2161 register_denial: None,
2162 delay_stream: None,
2163 }),
2164 2 => Ok(ScriptedSession {
2165 index: 2,
2166 log,
2167 events: vec![Ok(WorkerSessionEvent::Task(Box::new(second_task)))],
2168 fail_reports: false,
2169 register_denial: None,
2170 delay_stream: None,
2171 }),
2172 _ => Ok(ScriptedSession {
2173 index: attempt,
2174 log,
2175 events: Vec::new(),
2176 fail_reports: false,
2177 register_denial: Some(tonic::Status::permission_denied(
2178 "namespace `payments` revoked for subject `worker-a`",
2179 )),
2180 delay_stream: None,
2181 }),
2182 }
2183 }
2184 }
2185 };
2186
2187 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
2188 let result = tokio::time::timeout(Duration::from_secs(5), run)
2189 .await
2190 .map_err(WorkerError::decode)?;
2191
2192 drop(log_sender);
2193 let mut registrations = Vec::new();
2194 let mut reports = Vec::new();
2195 while let Some(entry) = log_receiver.recv().await {
2196 match entry {
2197 SessionLog::Registered(index, types) => registrations.push((index, types)),
2198 SessionLog::Reported(index, report) => reports.push((index, report)),
2199 }
2200 }
2201 assert_eq!(attempts.load(Ordering::SeqCst), 3);
2205 let expected_types = vec![String::from("double"), String::from("increment")];
2206 assert_eq!(
2207 registrations,
2208 vec![(1, expected_types.clone()), (2, expected_types)]
2209 );
2210 assert_eq!(reports.len(), 3);
2211 assert!(matches!(
2212 &reports[0],
2213 (1, RecordedReport::Completed(_, id, _)) if id == &first_activity
2214 ));
2215 assert!(matches!(
2216 &reports[1],
2217 (2, RecordedReport::Completed(_, id, _)) if id == &first_activity
2218 ));
2219 assert!(matches!(
2220 &reports[2],
2221 (2, RecordedReport::Completed(_, id, _)) if id == &second_activity
2222 ));
2223 let Err(error) = result else {
2224 return Err(WorkerError::decode(UnexpectedSuccess));
2225 };
2226 assert!(!error.is_retryable());
2227 assert!(matches!(
2228 error.grpc_status().map(tonic::Status::code),
2229 Some(tonic::Code::PermissionDenied)
2230 ));
2231 Ok(())
2232 }
2233
2234 #[tokio::test(start_paused = true)]
2235 async fn clean_close_loop_exhausts_drop_budget_with_classified_error() -> Result<(), WorkerError>
2236 {
2237 let worker = two_activity_worker()?;
2238 let attempts = Arc::new(AtomicUsize::new(0));
2239 let (log_sender, log_receiver) = mpsc::unbounded_channel();
2240 let connect = {
2241 let attempts = Arc::clone(&attempts);
2242 move || {
2243 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
2244 let log = log_sender.clone();
2245 async move {
2246 Ok(ScriptedSession {
2247 index: attempt,
2248 log,
2249 events: Vec::new(),
2250 fail_reports: false,
2251 register_denial: None,
2252 delay_stream: None,
2253 })
2254 }
2255 }
2256 };
2257
2258 let run = worker.run_with_connector_until(connect, std::future::pending::<()>());
2259 let result = tokio::time::timeout(Duration::from_secs(5), run)
2260 .await
2261 .map_err(WorkerError::decode)?;
2262
2263 assert_eq!(attempts.load(Ordering::SeqCst), 3);
2268 let Err(error) = result else {
2269 return Err(WorkerError::decode(UnexpectedSuccess));
2270 };
2271 assert!(matches!(error, WorkerError::CleanCloseExhausted));
2272 assert!(error.to_string().contains("closed the stream cleanly"));
2273 drop(log_receiver);
2274 Ok(())
2275 }
2276
2277 #[tokio::test]
2278 async fn shutdown_during_clean_close_backoff_returns_ok_promptly() -> Result<(), WorkerError> {
2279 let worker = two_activity_worker_with(slow_reconnect_config())?;
2280 let attempts = Arc::new(AtomicUsize::new(0));
2281 let (log_sender, log_receiver) = mpsc::unbounded_channel();
2282 let connect = {
2283 let attempts = Arc::clone(&attempts);
2284 move || {
2285 attempts.fetch_add(1, Ordering::SeqCst);
2286 let log = log_sender.clone();
2287 async move {
2288 Ok(ScriptedSession {
2289 index: 1,
2290 log,
2291 events: Vec::new(),
2292 fail_reports: false,
2293 register_denial: None,
2294 delay_stream: None,
2295 })
2296 }
2297 }
2298 };
2299 let shutdown = async {
2300 tokio::time::sleep(Duration::from_millis(50)).await;
2301 };
2302
2303 let run = worker.run_with_connector_until(connect, shutdown);
2306 tokio::time::timeout(Duration::from_millis(500), run)
2307 .await
2308 .map_err(WorkerError::decode)??;
2309
2310 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2311 drop(log_receiver);
2312 Ok(())
2313 }
2314
2315 #[tokio::test]
2319 async fn result_ack_clears_exactly_its_tracker_entry() -> Result<(), WorkerError> {
2320 use crate::protocol::reconnect::{PendingActivityReport, UnackedResultTracker};
2321 use crate::runtime::loop_::{SessionHealth, serve_activity_tasks_until};
2322
2323 let workflow_a = WorkflowId::new_v4();
2324 let workflow_b = WorkflowId::new_v4();
2325 let position = ActivityId::from_sequence_position(5);
2326 let mut tracker = UnackedResultTracker::new();
2327 for workflow in [&workflow_a, &workflow_b] {
2328 tracker.record(PendingActivityReport::Completed {
2329 workflow_id: workflow.clone(),
2330 activity_id: position.clone(),
2331 run_id: None,
2332 completion_token: String::from("generation-1"),
2333 output: Payload::new(ContentType::Json, b"{\"value\":1}".to_vec()),
2334 });
2335 }
2336
2337 let worker = two_activity_worker()?;
2338 let mut session = ChannelSession {
2339 receiver: None,
2340 reports: Vec::new(),
2341 registered: Vec::new(),
2342 };
2343 let (sender, receiver) = mpsc::channel(4);
2344 sender
2345 .send(Ok(WorkerSessionEvent::ResultAck {
2346 workflow_id: workflow_a.clone(),
2347 activity_id: position.clone(),
2348 }))
2349 .await
2350 .map_err(WorkerError::decode)?;
2351 sender
2353 .send(Ok(WorkerSessionEvent::ResultAck {
2354 workflow_id: WorkflowId::new_v4(),
2355 activity_id: ActivityId::from_sequence_position(99),
2356 }))
2357 .await
2358 .map_err(WorkerError::decode)?;
2359 drop(sender);
2360 session.receiver = Some(receiver);
2361
2362 let mut health = SessionHealth::default();
2363 serve_activity_tasks_until(
2364 &test_config(),
2365 &mut session,
2366 Arc::new(crate::activity::ActivityRegistry::new()),
2367 &mut tracker,
2368 &mut health,
2369 std::future::pending(),
2370 )
2371 .await?;
2372
2373 assert_eq!(tracker.len(), 1, "exactly the acked entry must clear");
2374 assert!(tracker.get(&workflow_a, &position).is_none());
2375 assert!(tracker.get(&workflow_b, &position).is_some());
2376 drop(worker);
2377 Ok(())
2378 }
2379
2380 #[tokio::test]
2384 async fn acked_results_decay_out_of_the_reconnect_replay() -> Result<(), WorkerError> {
2385 use crate::protocol::re_report_unacked;
2386 use crate::protocol::reconnect::{PendingActivityReport, UnackedResultTracker};
2387 use crate::runtime::loop_::{SessionHealth, serve_activity_tasks_until};
2388
2389 let workflow_id = WorkflowId::new_v4();
2390 let acked_id = ActivityId::from_sequence_position(1);
2391 let unacked_id = ActivityId::from_sequence_position(2);
2392 let mut tracker = UnackedResultTracker::new();
2393 for id in [&acked_id, &unacked_id] {
2394 tracker.record(PendingActivityReport::Completed {
2395 workflow_id: workflow_id.clone(),
2396 activity_id: id.clone(),
2397 run_id: None,
2398 completion_token: String::from("generation-1"),
2399 output: Payload::new(ContentType::Json, b"{\"value\":2}".to_vec()),
2400 });
2401 }
2402
2403 let mut session = ChannelSession {
2406 receiver: None,
2407 reports: Vec::new(),
2408 registered: Vec::new(),
2409 };
2410 let (sender, receiver) = mpsc::channel(2);
2411 sender
2412 .send(Ok(WorkerSessionEvent::ResultAck {
2413 workflow_id: workflow_id.clone(),
2414 activity_id: acked_id.clone(),
2415 }))
2416 .await
2417 .map_err(WorkerError::decode)?;
2418 drop(sender);
2419 session.receiver = Some(receiver);
2420 let mut health = SessionHealth::default();
2421 serve_activity_tasks_until(
2422 &test_config(),
2423 &mut session,
2424 Arc::new(crate::activity::ActivityRegistry::new()),
2425 &mut tracker,
2426 &mut health,
2427 std::future::pending(),
2428 )
2429 .await?;
2430
2431 let mut replay_session = ChannelSession {
2433 receiver: None,
2434 reports: Vec::new(),
2435 registered: Vec::new(),
2436 };
2437 re_report_unacked(&tracker, &mut replay_session).await?;
2438 assert_eq!(
2439 replay_session.reports.len(),
2440 1,
2441 "only the un-acked result may be re-reported"
2442 );
2443 assert!(matches!(
2444 &replay_session.reports[0],
2445 RecordedReport::Completed(_, id, _) if id == &unacked_id
2446 ));
2447
2448 let (sender, receiver) = mpsc::channel(2);
2451 sender
2452 .send(Ok(WorkerSessionEvent::ResultAck {
2453 workflow_id: workflow_id.clone(),
2454 activity_id: unacked_id.clone(),
2455 }))
2456 .await
2457 .map_err(WorkerError::decode)?;
2458 drop(sender);
2459 replay_session.receiver = Some(receiver);
2460 let mut health = SessionHealth::default();
2461 serve_activity_tasks_until(
2462 &test_config(),
2463 &mut replay_session,
2464 Arc::new(crate::activity::ActivityRegistry::new()),
2465 &mut tracker,
2466 &mut health,
2467 std::future::pending(),
2468 )
2469 .await?;
2470 assert!(tracker.is_empty(), "acks must drain the tracker");
2471
2472 let mut decayed_session = ChannelSession {
2473 receiver: None,
2474 reports: Vec::new(),
2475 registered: Vec::new(),
2476 };
2477 re_report_unacked(&tracker, &mut decayed_session).await?;
2478 assert!(
2479 decayed_session.reports.is_empty(),
2480 "steady-state replay must send nothing"
2481 );
2482 Ok(())
2483 }
2484
2485 #[tokio::test(start_paused = true)]
2488 async fn shutdown_interrupts_hung_unacked_replay_promptly() -> Result<(), WorkerError> {
2489 let workflow_id = WorkflowId::new_v4();
2492 let activity_id = ActivityId::from_sequence_position(3);
2493 let worker = two_activity_worker()?;
2494 let attempts = Arc::new(AtomicUsize::new(0));
2495 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
2496 let (registered_2_tx, registered_2_rx) = tokio::sync::oneshot::channel::<()>();
2497 let registered_2_tx = std::sync::Mutex::new(Some(registered_2_tx));
2498 let connect = {
2499 let log_sender = log_sender.clone();
2500 let workflow_id = workflow_id.clone();
2501 let activity_id = activity_id.clone();
2502 move |attempt_override: usize| {
2503 let log = log_sender.clone();
2504 let task = proto_task(workflow_id.clone(), activity_id.clone(), "double", 21);
2505 let notify = if attempt_override == 2 {
2506 registered_2_tx
2507 .lock()
2508 .ok()
2509 .and_then(|mut guard| guard.take())
2510 } else {
2511 None
2512 };
2513 async move {
2514 if attempt_override == 1 {
2515 Ok(SessionKind::Scripted(ScriptedSession {
2516 index: 1,
2517 log,
2518 events: vec![Ok(WorkerSessionEvent::Task(Box::new(task)))],
2519 fail_reports: true,
2520 register_denial: None,
2521 delay_stream: None,
2522 }))
2523 } else {
2524 if let Some(notify) = notify {
2525 let _ = notify.send(());
2526 }
2527 Ok(SessionKind::Hung(HungReportSession { index: 2, log }))
2528 }
2529 }
2530 }
2531 };
2532
2533 let attempts_for_connect = Arc::clone(&attempts);
2534 let run = worker.run_with_connector_until(
2535 move || {
2536 let attempt = attempts_for_connect.fetch_add(1, Ordering::SeqCst) + 1;
2537 connect(attempt)
2538 },
2539 async move {
2540 let _ = registered_2_rx.await;
2541 },
2542 );
2543
2544 tokio::time::timeout(Duration::from_secs(60), run)
2547 .await
2548 .map_err(WorkerError::decode)??;
2549
2550 drop(log_sender);
2551 let mut hung_session_reports = 0_usize;
2552 while let Some(entry) = log_receiver.recv().await {
2553 if let SessionLog::Reported(2, _) = entry {
2554 hung_session_reports += 1;
2555 }
2556 }
2557 assert_eq!(
2558 hung_session_reports, 0,
2559 "the hung replay must not have produced a report"
2560 );
2561 assert_eq!(attempts.load(Ordering::SeqCst), 2);
2562 Ok(())
2563 }
2564
2565 #[tokio::test(start_paused = true)]
2569 async fn drain_cycles_reconnect_without_consuming_drop_budget() -> Result<(), WorkerError> {
2570 let worker = two_activity_worker_with(test_config_with(ReconnectConfig::new(
2571 Duration::from_millis(5),
2572 Duration::from_millis(20),
2573 2,
2574 )))?;
2575 let attempts = Arc::new(AtomicUsize::new(0));
2576 let (log_sender, mut log_receiver) = mpsc::unbounded_channel();
2577 let connect = {
2578 let attempts = Arc::clone(&attempts);
2579 move || {
2580 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
2581 let log = log_sender.clone();
2582 async move {
2583 if attempt <= 3 {
2584 Ok(ScriptedSession {
2585 index: attempt,
2586 log,
2587 events: vec![Ok(WorkerSessionEvent::Drain)],
2588 fail_reports: false,
2589 register_denial: None,
2590 delay_stream: None,
2591 })
2592 } else {
2593 Ok(ScriptedSession {
2594 index: attempt,
2595 log,
2596 events: Vec::new(),
2597 fail_reports: false,
2598 register_denial: Some(tonic::Status::permission_denied(
2599 "namespace `payments` revoked for subject `worker-a`",
2600 )),
2601 delay_stream: None,
2602 })
2603 }
2604 }
2605 }
2606 };
2607
2608 let result = worker
2609 .run_with_connector_until(connect, std::future::pending::<()>())
2610 .await;
2611
2612 assert_eq!(attempts.load(Ordering::SeqCst), 4);
2616 let Err(error) = result else {
2617 return Err(WorkerError::decode(UnexpectedSuccess));
2618 };
2619 assert!(matches!(
2620 error.grpc_status().map(tonic::Status::code),
2621 Some(tonic::Code::PermissionDenied)
2622 ));
2623 let mut registrations = 0_usize;
2624 while let Some(entry) = log_receiver.recv().await {
2625 if matches!(entry, SessionLog::Registered(..)) {
2626 registrations += 1;
2627 }
2628 }
2629 assert_eq!(registrations, 3, "every drain cycle must re-register");
2630 Ok(())
2631 }
2632
2633 #[tokio::test(start_paused = true)]
2639 async fn drain_latch_keeps_abrupt_post_drain_failures_unbudgeted() -> Result<(), WorkerError> {
2640 let workflow_id = WorkflowId::new_v4();
2641 let worker = Worker::builder(test_config_with(ReconnectConfig::new(
2646 Duration::from_millis(5),
2647 Duration::from_millis(20),
2648 2,
2649 )))
2650 .register_activity("slow_double", |input: TestInput, context| {
2651 Box::pin(async move {
2652 let _ = context;
2653 tokio::time::sleep(Duration::from_millis(1)).await;
2654 Ok(TestOutput {
2655 value: input.value * 2,
2656 })
2657 })
2658 })?
2659 .build()?;
2660 let attempts = Arc::new(AtomicUsize::new(0));
2661 let (log_sender, log_receiver) = mpsc::unbounded_channel();
2662 let connect = {
2663 let attempts = Arc::clone(&attempts);
2664 let workflow_id = workflow_id.clone();
2665 move || {
2666 let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1;
2667 let log = log_sender.clone();
2668 let attempt_u64 = u64::try_from(attempt).unwrap_or(u64::MAX);
2669 let activity_id = ActivityId::from_sequence_position(attempt_u64);
2670 let task = proto_task(workflow_id.clone(), activity_id.clone(), "slow_double", 21);
2671 async move {
2672 if attempt <= 3 {
2673 Ok(LatchKind::Latch(DrainLatchSession {
2674 events: vec![
2675 Ok(WorkerSessionEvent::Task(Box::new(task))),
2676 Ok(WorkerSessionEvent::Drain),
2677 ],
2678 fail_id: activity_id,
2679 }))
2680 } else {
2681 Ok(LatchKind::Deny(ScriptedSession {
2682 index: attempt,
2683 log,
2684 events: Vec::new(),
2685 fail_reports: false,
2686 register_denial: Some(tonic::Status::permission_denied(
2687 "namespace `payments` revoked for subject `worker-a`",
2688 )),
2689 delay_stream: None,
2690 }))
2691 }
2692 }
2693 }
2694 };
2695
2696 let result = worker
2697 .run_with_connector_until(connect, std::future::pending::<()>())
2698 .await;
2699
2700 assert_eq!(attempts.load(Ordering::SeqCst), 4);
2703 let Err(error) = result else {
2704 return Err(WorkerError::decode(UnexpectedSuccess));
2705 };
2706 assert!(matches!(
2707 error.grpc_status().map(tonic::Status::code),
2708 Some(tonic::Code::PermissionDenied)
2709 ));
2710 drop(log_receiver);
2711 Ok(())
2712 }
2713
2714 #[tokio::test]
2718 async fn shutdown_during_post_drain_backoff_returns_ok_promptly() -> Result<(), WorkerError> {
2719 let worker = two_activity_worker_with(test_config_with(ReconnectConfig::new(
2720 Duration::from_secs(5),
2721 Duration::from_secs(10),
2722 5,
2723 )))?;
2724 let attempts = Arc::new(AtomicUsize::new(0));
2725 let (log_sender, log_receiver) = mpsc::unbounded_channel();
2726 let connect = {
2727 let attempts = Arc::clone(&attempts);
2728 move || {
2729 attempts.fetch_add(1, Ordering::SeqCst);
2730 let log = log_sender.clone();
2731 async move {
2732 Ok(ScriptedSession {
2733 index: 1,
2734 log,
2735 events: vec![Ok(WorkerSessionEvent::Drain)],
2736 fail_reports: false,
2737 register_denial: None,
2738 delay_stream: None,
2739 })
2740 }
2741 }
2742 };
2743 let shutdown = async {
2744 tokio::time::sleep(Duration::from_millis(50)).await;
2745 };
2746
2747 let run = worker.run_with_connector_until(connect, shutdown);
2750 tokio::time::timeout(Duration::from_millis(500), run)
2751 .await
2752 .map_err(WorkerError::decode)??;
2753
2754 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2755 drop(log_receiver);
2756 Ok(())
2757 }
2758}