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