1use std::collections::BTreeMap;
64use std::sync::{Arc, OnceLock};
65use std::time::{Duration, Instant};
66
67use aion::{ActivityDispatch, ActivityDispatcher};
68use aion_core::{ActivityId, ContentType, Payload, RunId, WorkflowId};
69use aion_proto::{ProtoActivityId, ProtoActivityTask, ProtoPayload, ProtoWorkflowId};
70use dashmap::DashMap;
71
72use super::dispatch::{ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink};
73use super::envelope::{CompletionFences, CompletionToken, idempotency_key};
74use super::heartbeat::{HeartbeatTracker, InFlightActivity};
75use super::queue_service::{
76 DeliveryRefusal, ExpiredClock, PARK_POLL_INTERVAL, PoolCensus, QueueDeclarationSource,
77 QueueServiceConfig, QueueServiceReason, QueueServiceState, SelectionRefusal, ServiceAddress,
78 ServiceWait, WorkerUnavailable, deliver_within_schedule_to_start, select_worker_or_refuse,
79};
80use super::registry::{
81 ConnectedWorkerRegistry, WorkerDelivery, WorkerHandle, WorkerId, WorkerMessage,
82};
83use crate::error::ServerError;
84use crate::shutdown::DrainState;
85use tracing::info_span;
86
87type SyncSender = std::sync::mpsc::SyncSender<Result<String, String>>;
88type SyncReceiver = std::sync::mpsc::Receiver<Result<String, String>>;
89
90type PendingActivityKey = (WorkflowId, ActivityId);
108
109pub trait OutboxDeliveryCallback: Send + Sync {
119 fn deliver_completion(
129 &self,
130 workflow_id: &WorkflowId,
131 activity_id: &ActivityId,
132 run_id: Option<&RunId>,
133 result: String,
134 ) -> Result<bool, ServerError>;
135
136 fn deliver_failure(
143 &self,
144 workflow_id: &WorkflowId,
145 activity_id: &ActivityId,
146 run_id: Option<&RunId>,
147 reason: String,
148 ) -> Result<bool, ServerError>;
149}
150
151#[derive(Clone, Default)]
163pub struct PendingActivities {
164 pending: Arc<DashMap<PendingActivityKey, SyncSender>>,
165 completion_fences: CompletionFences,
166 outbox_delivery: Arc<OnceLock<Arc<dyn OutboxDeliveryCallback>>>,
167 transport_losses: super::transport_loss::TransportLossLedger,
173}
174
175impl std::fmt::Debug for PendingActivities {
176 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
177 formatter
178 .debug_struct("PendingActivities")
179 .field("pending", &self.pending.len())
180 .field("completion_fences", &self.completion_fences)
181 .field(
182 "outbox_delivery_installed",
183 &self.outbox_delivery.get().is_some(),
184 )
185 .field("transport_losses", &self.transport_losses)
186 .finish()
187 }
188}
189
190impl PendingActivities {
191 fn insert(
192 &self,
193 workflow_id: WorkflowId,
194 activity_id: ActivityId,
195 ) -> Result<(CompletionToken, SyncReceiver), ServerError> {
196 let completion_token = self.completion_fences.issue(&workflow_id, &activity_id)?;
197 let (tx, rx) = std::sync::mpsc::sync_channel(1);
198 self.pending.insert((workflow_id, activity_id), tx);
199 Ok((completion_token, rx))
200 }
201
202 #[must_use]
204 pub fn completion_fences(&self) -> CompletionFences {
205 self.completion_fences.clone()
206 }
207
208 #[cfg(test)]
212 pub(crate) fn insert_for_test(
213 &self,
214 workflow_id: WorkflowId,
215 activity_id: ActivityId,
216 ) -> Result<(CompletionToken, SyncReceiver), ServerError> {
217 self.insert(workflow_id, activity_id)
218 }
219
220 pub fn set_outbox_delivery(&self, callback: Arc<dyn OutboxDeliveryCallback>) {
226 if self.outbox_delivery.set(callback).is_err() {
227 tracing::warn!("outbox delivery callback already installed; ignoring duplicate set");
228 }
229 }
230
231 fn complete(
239 &self,
240 workflow_id: &WorkflowId,
241 activity_id: &ActivityId,
242 run_id: Option<&RunId>,
243 result: Result<String, String>,
244 ) -> bool {
245 let matched = self
248 .pending
249 .remove(&(workflow_id.clone(), activity_id.clone()));
250 if let Some((_, sender)) = matched {
251 return sender.send(result).is_ok();
252 }
253 let Some(callback) = self.outbox_delivery.get() else {
254 return false;
256 };
257 let outcome = match result {
258 Ok(payload) => callback.deliver_completion(workflow_id, activity_id, run_id, payload),
259 Err(reason) => callback.deliver_failure(workflow_id, activity_id, run_id, reason),
260 };
261 match outcome {
262 Ok(true) => true,
263 Ok(false) => {
264 tracing::debug!(
266 workflow_id = %workflow_id,
267 activity_id = %activity_id,
268 "unmatched outbox completion for a workflow that is not currently live; \
269 recovery will re-arm it"
270 );
271 false
272 }
273 Err(error) => {
274 tracing::warn!(
275 workflow_id = %workflow_id,
276 activity_id = %activity_id,
277 %error,
278 "failed to deliver unmatched outbox completion to the live workflow"
279 );
280 false
281 }
282 }
283 }
284
285 fn complete_fenced(
288 &self,
289 workflow_id: &WorkflowId,
290 activity_id: &ActivityId,
291 run_id: Option<&RunId>,
292 completion_token: &CompletionToken,
293 result: Result<String, String>,
294 ) -> Result<bool, ServerError> {
295 self.completion_fences
296 .accept(workflow_id, activity_id, completion_token)
297 .inspect_err(|error| {
298 tracing::warn!(
299 workflow_id = %workflow_id,
300 activity_id = %activity_id,
301 %error,
302 "activity completion rejected by execution-generation fence"
303 );
304 })?;
305 let transport_domain = result
311 .as_ref()
312 .err()
313 .is_some_and(|reason| super::transport_loss::is_transport_domain_reason(reason));
314 if !transport_domain {
315 if let Err(error) = self.transport_losses.clear(workflow_id, activity_id) {
316 tracing::warn!(
317 workflow_id = %workflow_id,
318 activity_id = %activity_id,
319 %error,
320 "failed to retire the transport-loss budget for a resolved activity"
321 );
322 }
323 }
324 Ok(self.complete(workflow_id, activity_id, run_id, result))
325 }
326
327 #[must_use]
335 pub fn with_heartbeat_window(mut self, heartbeat_window: std::time::Duration) -> Self {
336 self.transport_losses = super::transport_loss::TransportLossLedger::new(heartbeat_window);
337 self
338 }
339
340 #[must_use]
342 pub const fn transport_losses(&self) -> &super::transport_loss::TransportLossLedger {
343 &self.transport_losses
344 }
345
346 fn classify_worker_loss(
354 &self,
355 workflow_id: &WorkflowId,
356 activity_id: &ActivityId,
357 worker_id: crate::worker::registry::WorkerId,
358 ) -> String {
359 let detail = super::transport_loss::worker_lost_detail(worker_id);
360 match self
361 .transport_losses
362 .record_loss(workflow_id, activity_id, &detail)
363 {
364 Ok(verdict) => {
365 if verdict.exhausted {
366 tracing::error!(
367 operation = "activity_complete",
368 workflow_id = %workflow_id,
369 activity_id = %activity_id,
370 worker_id = ?worker_id,
371 error_type = "TransportExhausted",
372 losses = verdict.losses,
373 budget_ms = self.transport_losses.budget().as_millis(),
374 "activity abandoned: the transport kept losing its worker past the \
375 transport-loss budget"
376 );
377 } else {
378 tracing::warn!(
379 operation = "activity_complete",
380 workflow_id = %workflow_id,
381 activity_id = %activity_id,
382 worker_id = ?worker_id,
383 error_type = "WorkerLost",
384 losses = verdict.losses,
385 budget_ms = self.transport_losses.budget().as_millis(),
386 "worker lost before reporting an activity result; the activity never ran \
387 and will be re-dispatched attempt-neutrally"
388 );
389 }
390 verdict.reason
391 }
392 Err(error) => {
393 tracing::error!(
394 workflow_id = %workflow_id,
395 activity_id = %activity_id,
396 %error,
397 "transport-loss ledger is unreadable; abandoning the activity rather than \
398 re-dispatching it without a budget"
399 );
400 format!(
401 "{}{detail} (transport-loss budget unreadable: {error})",
402 super::transport_loss::TRANSPORT_EXHAUSTED_REASON_PREFIX
403 )
404 }
405 }
406 }
407}
408
409impl ActivityCompletionSink for PendingActivities {
410 fn park_activity(
421 &self,
422 workflow_id: &WorkflowId,
423 activity_id: &ActivityId,
424 ) -> Result<(), ServerError> {
425 self.completion_fences
426 .revoke_current(workflow_id, activity_id)?;
427 let matched = self
428 .pending
429 .remove(&(workflow_id.clone(), activity_id.clone()));
430 if let Some((_, sender)) = matched {
431 let _ = sender.send(Err(aion::PARKED_ACTIVITY_REASON.to_owned()));
435 }
436 Ok(())
437 }
438
439 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
440 let result = match completion.outcome {
441 ActivityCompletionOutcome::Succeeded(payload) => {
442 payload_to_string(&payload).map_err(|reason| {
443 tracing::error!(
444 operation = "activity_complete",
445 workflow_id = %completion.workflow_id,
446 activity_id = %completion.activity_id,
447 error_type = "ActivityResultDecode",
448 %reason,
449 "activity completion failed"
450 );
451 ServerError::worker_dispatch("", "", format!("payload decode: {reason}"))
452 })?
453 }
454 ActivityCompletionOutcome::Failed(error) => {
455 let prefix = if error.is_retryable() {
456 "retryable"
457 } else {
458 "terminal"
459 };
460 tracing::error!(
461 operation = "activity_complete",
462 workflow_id = %completion.workflow_id,
463 activity_id = %completion.activity_id,
464 error_type = "ActivityFailed",
465 error_kind = prefix,
466 reason = %error.message,
467 "activity completion failed"
468 );
469 Err(format!("{prefix}:{}", error.message))
470 }
471 ActivityCompletionOutcome::WorkerLost { worker_id } => Err(self.classify_worker_loss(
476 &completion.workflow_id,
477 &completion.activity_id,
478 worker_id,
479 )),
480 };
481 self.complete_fenced(
482 &completion.workflow_id,
483 &completion.activity_id,
484 completion.run_id.as_ref(),
485 &completion.completion_token,
486 result,
487 )?;
488 Ok(())
489 }
490}
491
492fn payload_to_string(payload: &Payload) -> Result<Result<String, String>, String> {
493 match payload.content_type() {
494 ContentType::Json => String::from_utf8(payload.bytes().to_vec())
495 .map(Ok)
496 .map_err(|_| "activity result payload is not valid UTF-8".to_owned()),
497 }
498}
499
500pub struct WorkerActivityDispatcher {
508 registry: ConnectedWorkerRegistry,
509 namespace: String,
510 pending: PendingActivities,
511 heartbeat_tracker: HeartbeatTracker,
512 drain_state: DrainState,
513 tokio_handle: Option<tokio::runtime::Handle>,
514 attempt_owners: Option<super::intervention::AttemptOwnerIndex>,
520 queue_service: QueueServiceConfig,
524 queue_declarations: QueueDeclarationSource,
527 queue_state: QueueServiceState,
529}
530
531impl std::fmt::Debug for WorkerActivityDispatcher {
532 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
533 f.debug_struct("WorkerActivityDispatcher")
534 .field("namespace", &self.namespace)
535 .finish_non_exhaustive()
536 }
537}
538
539impl WorkerActivityDispatcher {
540 #[must_use]
548 pub fn new(
549 registry: ConnectedWorkerRegistry,
550 namespace: impl Into<String>,
551 heartbeat_tracker: HeartbeatTracker,
552 ) -> Self {
553 Self {
554 registry,
555 namespace: namespace.into(),
556 pending: PendingActivities::default(),
557 heartbeat_tracker,
558 drain_state: DrainState::default(),
559 tokio_handle: None,
560 attempt_owners: None,
561 queue_service: QueueServiceConfig::default(),
562 queue_declarations: QueueDeclarationSource::default(),
563 queue_state: QueueServiceState::default(),
564 }
565 }
566
567 #[must_use]
570 pub fn with_queue_service(mut self, queue_service: QueueServiceConfig) -> Self {
571 self.queue_service = queue_service;
572 self
573 }
574
575 #[must_use]
579 pub fn with_queue_declarations(mut self, queue_declarations: QueueDeclarationSource) -> Self {
580 self.queue_declarations = queue_declarations;
581 self
582 }
583
584 #[must_use]
587 pub fn with_queue_state(mut self, queue_state: QueueServiceState) -> Self {
588 self.queue_state = queue_state;
589 self
590 }
591
592 #[must_use]
599 pub fn with_attempt_owners(
600 mut self,
601 attempt_owners: super::intervention::AttemptOwnerIndex,
602 ) -> Self {
603 self.attempt_owners = Some(attempt_owners);
604 self
605 }
606
607 #[must_use]
609 pub fn with_pending(mut self, pending: PendingActivities) -> Self {
610 self.pending = pending;
611 self
612 }
613
614 #[must_use]
616 pub fn with_drain_state(mut self, drain_state: DrainState) -> Self {
617 self.drain_state = drain_state;
618 self
619 }
620
621 #[must_use]
623 pub fn with_tokio_handle(mut self, tokio_handle: tokio::runtime::Handle) -> Self {
624 self.tokio_handle = Some(tokio_handle);
625 self
626 }
627}
628
629impl WorkerActivityDispatcher {
630 fn ensure_accepting(
654 &self,
655 namespace: &str,
656 activity_type: &str,
657 workflow_id: &WorkflowId,
658 activity_id: &ActivityId,
659 worker_id: Option<WorkerId>,
660 ) -> Result<(), String> {
661 self.drain_state
662 .ensure_accepting(namespace, activity_type)
663 .map_err(|error| {
664 log_worker_error(
665 "WorkerDispatch",
666 namespace,
667 activity_type,
668 workflow_id,
669 activity_id,
670 worker_id,
671 &error.to_string(),
672 );
673 aion::PARKED_ACTIVITY_REASON.to_owned()
674 })
675 }
676
677 fn select_worker_or_wait(
685 &self,
686 address: &ServiceAddress,
687 workflow_id: &WorkflowId,
688 activity_id: &ActivityId,
689 ) -> Result<WorkerHandle, String> {
690 let wait = ServiceWait {
691 registry: &self.registry,
692 declarations: &self.queue_declarations,
693 config: &self.queue_service,
694 state: &self.queue_state,
695 address,
696 workflow_id,
697 activity_id,
698 };
699 let mut accepting = || {
700 self.ensure_accepting(
701 &address.namespace,
702 &address.activity_type,
703 workflow_id,
704 activity_id,
705 None,
706 )
707 };
708 let mut park = |budget: Option<Duration>| self.park_for_worker(budget);
709 select_worker_or_refuse(&wait, &mut accepting, &mut park).map_err(|refusal| {
710 let reason = refusal.reason_string();
711 if !matches!(refusal, SelectionRefusal::NotAccepting { .. }) {
714 let error_type = match refusal {
715 SelectionRefusal::Unavailable(_) => "WorkerUnavailable",
716 _ => "WorkerRegistry",
717 };
718 log_worker_error(
719 error_type,
720 &address.namespace,
721 &address.activity_type,
722 workflow_id,
723 activity_id,
724 None,
725 &reason,
726 );
727 }
728 reason
729 })
730 }
731
732 fn park_for_worker(&self, budget: Option<Duration>) {
760 let handle = self
761 .tokio_handle
762 .clone()
763 .or_else(|| tokio::runtime::Handle::try_current().ok());
764 let Some(handle) = handle else {
765 std::thread::sleep(
770 budget.map_or(PARK_POLL_INTERVAL, |budget| budget.min(PARK_POLL_INTERVAL)),
771 );
772 return;
773 };
774 handle.block_on(async {
775 let arrival = async {
776 tokio::select! {
777 () = self.registry.wait_for_worker() => {}
778 () = self.drain_state.wait_for_drain() => {}
779 }
780 };
781 match budget {
782 None => arrival.await,
783 Some(budget) => {
784 drop(tokio::time::timeout(budget, arrival).await);
787 }
788 }
789 });
790 }
791
792 fn track_worker_task(
793 &self,
794 worker_id: WorkerId,
795 activity_type: &str,
796 workflow_id: &WorkflowId,
797 activity_id: &ActivityId,
798 completion_token: CompletionToken,
799 ) -> Result<(), String> {
800 self.heartbeat_tracker
801 .track_task(
802 worker_id,
803 InFlightActivity {
804 workflow_id: workflow_id.clone(),
805 activity_id: activity_id.clone(),
806 completion_token,
807 },
808 Instant::now(),
809 )
810 .map_err(|error| {
811 let reason = error.to_string();
812 log_worker_error(
813 "WorkerHeartbeatTracker",
814 &self.namespace,
815 activity_type,
816 workflow_id,
817 activity_id,
818 Some(worker_id),
819 &reason,
820 );
821 reason
822 })
823 }
824
825 fn cleanup_activity(
826 &self,
827 worker_id: WorkerId,
828 workflow_id: &WorkflowId,
829 activity_id: &ActivityId,
830 completion_token: &CompletionToken,
831 ) {
832 self.pending
833 .pending
834 .remove(&(workflow_id.clone(), activity_id.clone()));
835 if let Err(error) =
836 self.pending
837 .completion_fences
838 .revoke(workflow_id, activity_id, completion_token)
839 {
840 tracing::error!(
841 workflow_id = %workflow_id,
842 activity_id = %activity_id,
843 %error,
844 "failed to revoke undelivered activity generation"
845 );
846 }
847 let _ = self
848 .heartbeat_tracker
849 .complete_task(worker_id, workflow_id, activity_id);
850 self.drain_state.notify_activity_drained();
851 }
852
853 fn send_activity_task(
863 &self,
864 worker: &WorkerHandle,
865 task: ProtoActivityTask,
866 address: &ServiceAddress,
867 workflow_id: &WorkflowId,
868 activity_id: &ActivityId,
869 completion_token: &CompletionToken,
870 ) -> Result<(), String> {
871 match worker.delivery() {
872 WorkerDelivery::Grpc(sender) => {
873 let worker_id = worker.id();
874 let mut accepting = || {
875 self.ensure_accepting(
876 &address.namespace,
877 &address.activity_type,
878 workflow_id,
879 activity_id,
880 Some(worker_id),
881 )
882 };
883 let handed_over = deliver_within_schedule_to_start(
884 sender,
885 WorkerMessage::ActivityTask(Box::new(task)),
886 self.queue_service.schedule_to_start_timeout,
887 &mut accepting,
888 );
889 let Err(refusal) = handed_over else {
890 return Ok(());
891 };
892 self.cleanup_activity(worker_id, workflow_id, activity_id, completion_token);
893 let (error_type, reason) = self.hand_off_failure(&refusal, address);
894 if !matches!(refusal, DeliveryRefusal::NotAccepting { .. }) {
895 log_worker_error(
896 error_type,
897 &address.namespace,
898 &address.activity_type,
899 workflow_id,
900 activity_id,
901 Some(worker_id),
902 &reason,
903 );
904 }
905 Err(reason)
906 }
907 #[cfg(feature = "liminal-transport")]
914 WorkerDelivery::Liminal(delivery) => self.send_liminal_activity_task(
915 worker.id(),
916 delivery,
917 task,
918 &address.activity_type,
919 workflow_id,
920 activity_id,
921 ),
922 }
923 }
924
925 fn hand_off_failure(
932 &self,
933 refusal: &DeliveryRefusal,
934 address: &ServiceAddress,
935 ) -> (&'static str, String) {
936 match refusal {
937 DeliveryRefusal::Saturated { waited } => {
938 let census = self
941 .registry
942 .pool_census(
943 &address.namespace,
944 &address.task_queue,
945 &address.activity_type,
946 address.node.as_deref(),
947 )
948 .unwrap_or_else(|error| {
949 tracing::error!(
950 namespace = %address.namespace,
951 task_queue = %address.task_queue,
952 activity_type = %address.activity_type,
953 %error,
954 "poller census failed while reporting a saturated queue; \
955 the refusal carries an empty census"
956 );
957 PoolCensus::default()
958 });
959 let unavailable = WorkerUnavailable {
960 reason: QueueServiceReason::Saturated,
961 clock: Some(ExpiredClock::ScheduleToStart),
962 waited: *waited,
963 address: address.clone(),
964 census,
965 };
966 ("WorkerUnavailable", unavailable.reason_string())
967 }
968 DeliveryRefusal::Full => (
969 "WorkerChannelClosed",
970 "worker task channel full or closed: no available capacity".to_owned(),
971 ),
972 DeliveryRefusal::Closed => (
973 "WorkerChannelClosed",
974 "worker task channel full or closed: channel closed".to_owned(),
975 ),
976 DeliveryRefusal::NotAccepting { reason } => ("WorkerDispatch", reason.clone()),
977 }
978 }
979
980 #[cfg(feature = "liminal-transport")]
1004 fn send_liminal_activity_task(
1005 &self,
1006 worker_id: WorkerId,
1007 delivery: &super::liminal_transport::LiminalWorkerDelivery,
1008 task: ProtoActivityTask,
1009 activity_type: &str,
1010 workflow_id: &WorkflowId,
1011 activity_id: &ActivityId,
1012 ) -> Result<(), String> {
1013 let completion_token =
1014 CompletionToken::from_wire(workflow_id, activity_id, task.completion_token.clone())
1015 .map_err(|error| error.to_string())?;
1016 let heartbeat_window_ms =
1017 u64::try_from(self.heartbeat_tracker.heartbeat_window().as_millis())
1018 .unwrap_or(u64::MAX);
1019 let attempt = task.attempt;
1020 let request = super::liminal_transport::DispatchRequest {
1021 activity_type: activity_type.to_owned(),
1022 workflow_id: workflow_id.clone(),
1023 ordinal: activity_id.sequence_position(),
1024 run_id: task
1025 .run_id
1026 .map(RunId::try_from)
1027 .transpose()
1028 .map_err(|error| error.to_string())?,
1029 attempt,
1030 completion_token: task.completion_token,
1031 idempotency_key: task.idempotency_key,
1032 labels: task.labels.into_iter().collect(),
1033 heartbeat_window_ms,
1034 input: task.input.map(|payload| payload.bytes).unwrap_or_default(),
1035 };
1036 let awaiter = match delivery.push_dispatch(&request) {
1040 Ok(awaiter) => awaiter,
1041 Err(error) => {
1042 let reason = format!("worker liminal push failed: {error}");
1043 self.cleanup_activity(worker_id, workflow_id, activity_id, &completion_token);
1044 log_worker_error(
1045 "WorkerChannelClosed",
1046 &self.namespace,
1047 activity_type,
1048 workflow_id,
1049 activity_id,
1050 Some(worker_id),
1051 &reason,
1052 );
1053 return Err(reason);
1054 }
1055 };
1056 let owner_binding = self.attempt_owners.as_ref().map(|owners| {
1060 super::liminal_transport::AttemptOwnerGuard::bind(
1061 owners.clone(),
1062 super::intervention::AttemptKey::new(
1063 workflow_id.clone(),
1064 activity_id.clone(),
1065 attempt,
1066 ),
1067 worker_id,
1068 )
1069 });
1070 self.spawn_liminal_reply_router(
1071 worker_id,
1072 awaiter,
1073 workflow_id,
1074 activity_id,
1075 &completion_token,
1076 owner_binding,
1077 );
1078 Ok(())
1079 }
1080
1081 #[cfg(feature = "liminal-transport")]
1121 fn spawn_liminal_reply_router(
1122 &self,
1123 worker_id: WorkerId,
1124 awaiter: liminal_server::server::connection::PushReplyAwaiter,
1125 workflow_id: &WorkflowId,
1126 activity_id: &ActivityId,
1127 completion_token: &CompletionToken,
1128 owner_binding: Option<super::liminal_transport::AttemptOwnerGuard>,
1129 ) {
1130 let pending = self.pending.clone();
1131 let heartbeat_tracker = self.heartbeat_tracker.clone();
1132 let drain_state = self.drain_state.clone();
1133 let workflow_id = workflow_id.clone();
1134 let activity_id = activity_id.clone();
1135 let completion_token = completion_token.clone();
1136 std::thread::spawn(move || {
1137 let _owner_binding = owner_binding;
1141 route_liminal_reply(
1142 &pending,
1143 &heartbeat_tracker,
1144 &drain_state,
1145 &awaiter,
1146 (worker_id, &workflow_id, &activity_id, &completion_token),
1147 );
1148 });
1149 }
1150
1151 fn await_activity_result(
1155 &self,
1156 context: &ActivityDispatchContext<'_>,
1157 rx: &SyncReceiver,
1158 ) -> Result<String, String> {
1159 match self.registry.is_registered(context.worker_id) {
1167 Ok(true) => {}
1168 Ok(false) => {
1169 if let Ok(result) = rx.try_recv() {
1173 return self.deliver_result(context, result);
1174 }
1175 self.cleanup_activity(
1176 context.worker_id,
1177 context.workflow_id,
1178 context.activity_id,
1179 &context.completion_token,
1180 );
1181 let reason = self.pending.classify_worker_loss(
1186 context.workflow_id,
1187 context.activity_id,
1188 context.worker_id,
1189 );
1190 log_worker_error(
1191 "WorkerLost",
1192 &self.namespace,
1193 context.activity_type,
1194 context.workflow_id,
1195 context.activity_id,
1196 Some(context.worker_id),
1197 &reason,
1198 );
1199 return Err(reason);
1200 }
1201 Err(error) => {
1202 self.cleanup_activity(
1203 context.worker_id,
1204 context.workflow_id,
1205 context.activity_id,
1206 &context.completion_token,
1207 );
1208 let reason = format!("worker registry inspection failed: {error}");
1209 log_worker_error(
1210 "WorkerRegistry",
1211 &self.namespace,
1212 context.activity_type,
1213 context.workflow_id,
1214 context.activity_id,
1215 Some(context.worker_id),
1216 &reason,
1217 );
1218 return Err(reason);
1219 }
1220 }
1221 if let Ok(result) = rx.recv() {
1222 return self.deliver_result(context, result);
1223 }
1224 self.cleanup_activity(
1227 context.worker_id,
1228 context.workflow_id,
1229 context.activity_id,
1230 &context.completion_token,
1231 );
1232 let reason = "activity response channel dropped".to_owned();
1233 log_worker_error(
1234 "WorkerChannelClosed",
1235 &self.namespace,
1236 context.activity_type,
1237 context.workflow_id,
1238 context.activity_id,
1239 Some(context.worker_id),
1240 &reason,
1241 );
1242 Err(reason)
1243 }
1244
1245 fn deliver_result(
1246 &self,
1247 context: &ActivityDispatchContext<'_>,
1248 result: Result<String, String>,
1249 ) -> Result<String, String> {
1250 self.pending
1251 .pending
1252 .remove(&(context.workflow_id.clone(), context.activity_id.clone()));
1253 if let Err(reason) = &result
1258 && aion::is_parked_reason(reason)
1259 {
1260 tracing::info!(
1261 operation = "activity_dispatch",
1262 namespace = %self.namespace,
1263 workflow_id = %context.workflow_id,
1264 activity_id = %context.activity_id,
1265 activity_type = context.activity_type,
1266 worker_id = ?context.worker_id,
1267 "activity parked for restart recovery"
1268 );
1269 return result;
1270 }
1271 log_activity_completion(context, result.is_ok());
1272 result.inspect_err(|reason| {
1273 log_worker_error(
1274 "ActivityFailed",
1275 &self.namespace,
1276 context.activity_type,
1277 context.workflow_id,
1278 context.activity_id,
1279 Some(context.worker_id),
1280 reason,
1281 );
1282 })
1283 }
1284}
1285
1286impl ActivityDispatcher for WorkerActivityDispatcher {
1287 fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
1288 match tokio::runtime::Handle::try_current() {
1289 Ok(handle) => match handle.runtime_flavor() {
1290 tokio::runtime::RuntimeFlavor::MultiThread => {
1291 tokio::task::block_in_place(|| self.dispatch_blocking(request))
1298 }
1299 flavor => Err(format!(
1300 "activity dispatch blocks the calling thread until the worker responds; \
1301 a {flavor:?} tokio runtime cannot host that wait because the worker \
1302 stream forwarder shares its only executor thread and the task could \
1303 never be delivered — run the engine on a multi-thread tokio runtime"
1304 )),
1305 },
1306 Err(_) => self.dispatch_blocking(request),
1310 }
1311 }
1312}
1313
1314impl WorkerActivityDispatcher {
1315 fn dispatch_blocking(&self, request: ActivityDispatch) -> Result<String, String> {
1332 let ActivityDispatch {
1333 namespace,
1334 task_queue,
1335 node,
1339 workflow_id,
1340 run_id,
1341 activity_id,
1342 name,
1343 input,
1344 config: _,
1345 attempt,
1346 labels,
1347 advisory: _,
1352 } = request;
1353 let started_at = Instant::now();
1354 self.ensure_accepting(&namespace, &name, &workflow_id, &activity_id, None)?;
1355 let address = ServiceAddress {
1356 namespace: namespace.clone(),
1357 task_queue: task_queue.clone(),
1358 activity_type: name.clone(),
1359 node: node.clone(),
1360 };
1361 let worker = self.select_worker_or_wait(&address, &workflow_id, &activity_id)?;
1362 let worker_id = worker.id();
1363 let span = info_span!(
1364 "activity_dispatch",
1365 operation = "activity_dispatch",
1366 namespace = %namespace,
1367 task_queue = %task_queue,
1368 node = node.as_deref(),
1369 workflow_id = %workflow_id,
1370 activity_id = %activity_id,
1371 activity_type = %name,
1372 worker_id = ?worker_id,
1373 );
1374 let _span_guard = span.enter();
1375 self.ensure_accepting(
1376 &namespace,
1377 &name,
1378 &workflow_id,
1379 &activity_id,
1380 Some(worker_id),
1381 )?;
1382
1383 let (completion_token, rx) = self
1384 .pending
1385 .insert(workflow_id.clone(), activity_id.clone())
1386 .map_err(|error| error.to_string())?;
1387 let task = activity_task(
1388 &name,
1389 &input,
1390 (&workflow_id, &run_id, &activity_id),
1391 attempt,
1392 labels,
1393 &completion_token,
1394 );
1395 if let Err(error) = self.track_worker_task(
1396 worker_id,
1397 &name,
1398 &workflow_id,
1399 &activity_id,
1400 completion_token.clone(),
1401 ) {
1402 self.cleanup_activity(worker_id, &workflow_id, &activity_id, &completion_token);
1403 return Err(error);
1404 }
1405 self.send_activity_task(
1406 &worker,
1407 task,
1408 &address,
1409 &workflow_id,
1410 &activity_id,
1411 &completion_token,
1412 )?;
1413 let context = ActivityDispatchContext {
1414 namespace: &namespace,
1415 activity_type: &name,
1416 worker_id,
1417 workflow_id: &workflow_id,
1418 activity_id: &activity_id,
1419 completion_token,
1420 started_at,
1421 };
1422 self.await_activity_result(&context, &rx)
1423 }
1424}
1425
1426#[cfg(feature = "liminal-transport")]
1433fn route_liminal_reply(
1434 pending: &PendingActivities,
1435 heartbeat_tracker: &HeartbeatTracker,
1436 drain_state: &DrainState,
1437 awaiter: &liminal_server::server::connection::PushReplyAwaiter,
1438 execution: (WorkerId, &WorkflowId, &ActivityId, &CompletionToken),
1439) {
1440 let (worker_id, workflow_id, activity_id, current_token) = execution;
1441 let waited = super::liminal_transport::receive_bridge_reply(awaiter, || {
1446 heartbeat_tracker
1447 .is_tracked(worker_id, workflow_id, activity_id)
1448 .unwrap_or(false)
1449 });
1450 let (run_id, submitted_token, outcome, synthesized) = match waited {
1454 Ok(Some(response)) => {
1455 let submitted_token = match CompletionToken::from_wire(
1456 workflow_id,
1457 activity_id,
1458 response.completion_token,
1459 ) {
1460 Ok(token) => token,
1461 Err(error) => {
1462 tracing::warn!(
1463 worker_id = ?worker_id,
1464 workflow_id = %workflow_id,
1465 activity_id = %activity_id,
1466 %error,
1467 "liminal activity completion omitted its generation proof"
1468 );
1469 return;
1470 }
1471 };
1472 (response.run_id, submitted_token, response.outcome, false)
1473 }
1474 Ok(None) => {
1475 tracing::debug!(
1476 worker_id = ?worker_id,
1477 workflow_id = %workflow_id,
1478 activity_id = %activity_id,
1479 "liminal dispatch resolved by another path; abandoning reply wait"
1480 );
1481 return;
1482 }
1483 Err(error) if error.is_worker_connection_lost() => (
1484 None,
1485 current_token.clone(),
1486 Err(pending.classify_worker_loss(workflow_id, activity_id, worker_id)),
1490 true,
1491 ),
1492 Err(error) => (
1493 None,
1494 current_token.clone(),
1495 Err(format!("retryable:worker liminal reply failed: {error}")),
1496 true,
1497 ),
1498 };
1499 if synthesized {
1504 let was_tracked =
1505 complete_liminal_tracking(heartbeat_tracker, worker_id, workflow_id, activity_id);
1506 if !was_tracked {
1507 tracing::debug!(
1510 worker_id = ?worker_id,
1511 workflow_id = %workflow_id,
1512 activity_id = %activity_id,
1513 "liminal dispatch already resolved; dropping synthesized lost-worker failure"
1514 );
1515 return;
1516 }
1517 }
1518 if let Err(error) = pending.complete_fenced(
1519 workflow_id,
1520 activity_id,
1521 run_id.as_ref(),
1522 &submitted_token,
1523 outcome,
1524 ) {
1525 tracing::warn!(
1526 worker_id = ?worker_id,
1527 workflow_id = %workflow_id,
1528 activity_id = %activity_id,
1529 %error,
1530 "liminal activity completion handoff rejected"
1531 );
1532 return;
1533 }
1534 if !synthesized
1535 && let Err(error) = heartbeat_tracker.complete_task(worker_id, workflow_id, activity_id)
1536 {
1537 tracing::error!(
1538 worker_id = ?worker_id,
1539 workflow_id = %workflow_id,
1540 activity_id = %activity_id,
1541 %error,
1542 "failed to clear in-flight tracking for completed liminal activity"
1543 );
1544 }
1545 drain_state.notify_activity_drained();
1546}
1547
1548#[cfg(feature = "liminal-transport")]
1549fn complete_liminal_tracking(
1550 heartbeat_tracker: &HeartbeatTracker,
1551 worker_id: WorkerId,
1552 workflow_id: &WorkflowId,
1553 activity_id: &ActivityId,
1554) -> bool {
1555 heartbeat_tracker
1556 .complete_task(worker_id, workflow_id, activity_id)
1557 .unwrap_or_else(|error| {
1558 tracing::error!(
1559 worker_id = ?worker_id,
1560 workflow_id = %workflow_id,
1561 activity_id = %activity_id,
1562 %error,
1563 "failed to clear in-flight tracking for completed liminal activity"
1564 );
1565 true
1566 })
1567}
1568
1569struct ActivityDispatchContext<'a> {
1570 namespace: &'a str,
1571 activity_type: &'a str,
1572 worker_id: WorkerId,
1573 workflow_id: &'a WorkflowId,
1574 activity_id: &'a ActivityId,
1575 completion_token: CompletionToken,
1576 started_at: Instant,
1577}
1578
1579fn activity_task(
1580 activity_type: &str,
1581 input: &str,
1582 execution: (&WorkflowId, &RunId, &ActivityId),
1583 attempt: u32,
1584 labels: BTreeMap<String, String>,
1585 completion_token: &CompletionToken,
1586) -> ProtoActivityTask {
1587 let (workflow_id, run_id, activity_id) = execution;
1588 ProtoActivityTask {
1589 workflow_id: Some(ProtoWorkflowId::from(workflow_id.clone())),
1590 activity_id: Some(ProtoActivityId::from(activity_id.clone())),
1591 activity_type: activity_type.to_owned(),
1592 input: Some(ProtoPayload {
1593 content_type: String::from("application/json"),
1594 bytes: input.as_bytes().to_vec(),
1595 }),
1596 attempt,
1597 labels: labels.into_iter().collect(),
1598 run_id: Some(run_id.clone().into()),
1599 completion_token: completion_token.as_str().to_owned(),
1600 idempotency_key: idempotency_key(workflow_id, run_id, activity_id),
1601 }
1602}
1603
1604fn log_activity_completion(context: &ActivityDispatchContext<'_>, succeeded: bool) {
1605 let duration_ms = duration_ms(context.started_at.elapsed());
1606 tracing::info!(
1607 operation = "activity_complete",
1608 namespace = context.namespace,
1609 workflow_id = %context.workflow_id,
1610 activity_id = %context.activity_id,
1611 activity_type = context.activity_type,
1612 worker_id = ?context.worker_id,
1613 duration_ms,
1614 outcome = if succeeded { "succeeded" } else { "failed" },
1615 "activity completed"
1616 );
1617}
1618
1619fn duration_ms(duration: Duration) -> u64 {
1620 u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
1621}
1622
1623fn log_worker_error(
1624 error_type: &'static str,
1625 namespace: &str,
1626 activity_type: &str,
1627 workflow_id: &WorkflowId,
1628 activity_id: &ActivityId,
1629 worker_id: Option<super::registry::WorkerId>,
1630 reason: &str,
1631) {
1632 tracing::error!(
1633 operation = "activity_dispatch",
1634 namespace,
1635 workflow_id = %workflow_id,
1636 activity_id = %activity_id,
1637 activity_type,
1638 worker_id = ?worker_id,
1639 error_type,
1640 reason,
1641 "worker interaction failed"
1642 );
1643}
1644
1645#[cfg(test)]
1646mod tests {
1647 use std::sync::Mutex;
1648
1649 use aion_core::{ActivityError, ActivityErrorKind, ContentType, Payload};
1650
1651 use super::*;
1652
1653 fn activity_id(pos: u64) -> ActivityId {
1654 ActivityId::from_sequence_position(pos)
1655 }
1656
1657 #[test]
1658 fn pending_insert_and_complete_delivers_result() -> Result<(), ServerError> {
1659 let pending = PendingActivities::default();
1660 let workflow_id = WorkflowId::new_v4();
1661 let id = activity_id(1);
1662 let rx = pending.insert(workflow_id.clone(), id.clone())?.1;
1663
1664 assert!(pending.complete(&workflow_id, &id, None, Ok("done".to_owned())));
1665 assert_eq!(
1666 rx.recv_timeout(Duration::from_millis(50)),
1667 Ok(Ok("done".to_owned()))
1668 );
1669 Ok(())
1670 }
1671
1672 #[test]
1673 fn pending_complete_unknown_returns_false() {
1674 let pending = PendingActivities::default();
1675 assert!(!pending.complete(
1676 &WorkflowId::new_v4(),
1677 &activity_id(99),
1678 None,
1679 Ok("orphan".to_owned())
1680 ));
1681 }
1682
1683 #[derive(Default)]
1684 struct RecordingOutboxCallback {
1685 completions: Mutex<Vec<(WorkflowId, ActivityId, String)>>,
1686 failures: Mutex<Vec<(WorkflowId, ActivityId, String)>>,
1687 live: bool,
1688 }
1689
1690 impl OutboxDeliveryCallback for RecordingOutboxCallback {
1691 fn deliver_completion(
1692 &self,
1693 workflow_id: &WorkflowId,
1694 activity_id: &ActivityId,
1695 run_id: Option<&RunId>,
1696 result: String,
1697 ) -> Result<bool, ServerError> {
1698 let _ = run_id;
1699 self.completions
1700 .lock()
1701 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1702 .push((workflow_id.clone(), activity_id.clone(), result));
1703 Ok(self.live)
1704 }
1705
1706 fn deliver_failure(
1707 &self,
1708 workflow_id: &WorkflowId,
1709 activity_id: &ActivityId,
1710 run_id: Option<&RunId>,
1711 reason: String,
1712 ) -> Result<bool, ServerError> {
1713 let _ = run_id;
1714 self.failures
1715 .lock()
1716 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1717 .push((workflow_id.clone(), activity_id.clone(), reason));
1718 Ok(self.live)
1719 }
1720 }
1721
1722 #[test]
1723 fn unmatched_completion_routes_to_outbox_callback_when_installed() -> Result<(), ServerError> {
1724 let pending = PendingActivities::default();
1725 let callback = Arc::new(RecordingOutboxCallback {
1726 live: true,
1727 ..RecordingOutboxCallback::default()
1728 });
1729 pending.clone().set_outbox_delivery(callback.clone());
1731
1732 let workflow_id = WorkflowId::new_v4();
1733 let id = activity_id(7);
1734
1735 assert!(pending.complete(&workflow_id, &id, None, Ok("done".to_owned())));
1738 let completions = callback
1739 .completions
1740 .lock()
1741 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?;
1742 assert_eq!(completions.len(), 1);
1743 assert_eq!(completions[0].0, workflow_id);
1744 assert_eq!(completions[0].1, id);
1745 assert_eq!(completions[0].2, "done");
1746 Ok(())
1747 }
1748
1749 #[test]
1750 fn unmatched_failure_routes_to_outbox_callback_and_not_live_reports_false()
1751 -> Result<(), ServerError> {
1752 let pending = PendingActivities::default();
1753 let callback = Arc::new(RecordingOutboxCallback::default());
1755 pending.set_outbox_delivery(callback.clone());
1756
1757 let workflow_id = WorkflowId::new_v4();
1758 let id = activity_id(8);
1759
1760 assert!(!pending.complete(&workflow_id, &id, None, Err("retryable:boom".to_owned())));
1761 let failures = callback
1762 .failures
1763 .lock()
1764 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?;
1765 assert_eq!(failures.len(), 1);
1766 assert_eq!(failures[0].2, "retryable:boom");
1767 Ok(())
1768 }
1769
1770 #[test]
1771 fn unmatched_completion_is_silent_drop_when_no_callback_installed() {
1772 let pending = PendingActivities::default();
1775 assert!(!pending.complete(
1776 &WorkflowId::new_v4(),
1777 &activity_id(9),
1778 None,
1779 Ok("x".to_owned())
1780 ));
1781 }
1782
1783 #[test]
1784 fn matched_completion_never_reaches_outbox_callback() -> Result<(), ServerError> {
1785 let pending = PendingActivities::default();
1786 let callback = Arc::new(RecordingOutboxCallback {
1787 live: true,
1788 ..RecordingOutboxCallback::default()
1789 });
1790 pending.set_outbox_delivery(callback.clone());
1791
1792 let workflow_id = WorkflowId::new_v4();
1793 let id = activity_id(10);
1794 let rx = pending.insert(workflow_id.clone(), id.clone())?.1;
1795
1796 assert!(pending.complete(&workflow_id, &id, None, Ok("matched".to_owned())));
1797 assert_eq!(
1798 rx.recv_timeout(Duration::from_millis(50)),
1799 Ok(Ok("matched".to_owned()))
1800 );
1801 assert!(
1802 callback
1803 .completions
1804 .lock()
1805 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1806 .is_empty(),
1807 "a matched completion must deliver to its waiter, not the outbox callback"
1808 );
1809 Ok(())
1810 }
1811
1812 #[test]
1816 fn park_activity_resolves_matched_waiter_with_the_parked_sentinel() -> Result<(), ServerError> {
1817 let pending = PendingActivities::default();
1818 let workflow_id = WorkflowId::new_v4();
1819 let id = activity_id(11);
1820 let rx = pending.insert(workflow_id.clone(), id.clone())?.1;
1821
1822 pending.park_activity(&workflow_id, &id)?;
1823 let result = rx
1824 .recv_timeout(Duration::from_millis(50))
1825 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1826 assert_eq!(result, Err(aion::PARKED_ACTIVITY_REASON.to_owned()));
1827 Ok(())
1828 }
1829
1830 #[test]
1834 fn unmatched_park_is_a_noop_and_never_reaches_the_outbox_callback() -> Result<(), ServerError> {
1835 let pending = PendingActivities::default();
1836 let callback = Arc::new(RecordingOutboxCallback {
1837 live: true,
1838 ..RecordingOutboxCallback::default()
1839 });
1840 pending.set_outbox_delivery(callback.clone());
1841
1842 pending.park_activity(&WorkflowId::new_v4(), &activity_id(12))?;
1843
1844 assert!(
1845 callback
1846 .failures
1847 .lock()
1848 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1849 .is_empty(),
1850 "a park must never be delivered as an outbox failure"
1851 );
1852 assert!(
1853 callback
1854 .completions
1855 .lock()
1856 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1857 .is_empty(),
1858 "a park must never be delivered as an outbox completion"
1859 );
1860 Ok(())
1861 }
1862
1863 #[test]
1864 fn completion_sink_routes_success() -> Result<(), ServerError> {
1865 let pending = PendingActivities::default();
1866 let workflow_id = WorkflowId::new_v4();
1867 let id = activity_id(2);
1868 let (completion_token, rx) = pending.insert(workflow_id.clone(), id.clone())?;
1869 let payload = Payload::new(ContentType::Json, br#"{"greeting":"hi"}"#.to_vec());
1870
1871 pending.complete_activity(ActivityCompletion {
1872 workflow_id,
1873 activity_id: id,
1874 run_id: None,
1875 completion_token,
1876 outcome: ActivityCompletionOutcome::Succeeded(payload),
1877 })?;
1878
1879 let result = rx
1880 .recv_timeout(Duration::from_millis(50))
1881 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1882 assert_eq!(result, Ok(r#"{"greeting":"hi"}"#.to_owned()));
1883 Ok(())
1884 }
1885
1886 #[test]
1887 fn malformed_payload_does_not_consume_the_current_generation() -> Result<(), ServerError> {
1888 let pending = PendingActivities::default();
1889 let workflow_id = WorkflowId::new_v4();
1890 let id = activity_id(12);
1891 let (completion_token, rx) = pending.insert(workflow_id.clone(), id.clone())?;
1892
1893 let malformed = pending.complete_activity(ActivityCompletion {
1894 workflow_id: workflow_id.clone(),
1895 activity_id: id.clone(),
1896 run_id: None,
1897 completion_token: completion_token.clone(),
1898 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1899 ContentType::Json,
1900 vec![0xff],
1901 )),
1902 });
1903 assert!(matches!(malformed, Err(ServerError::WorkerDispatch { .. })));
1904 assert!(
1905 rx.try_recv().is_err(),
1906 "an invalid result must leave the waiter unresolved"
1907 );
1908
1909 pending.complete_activity(ActivityCompletion {
1910 workflow_id,
1911 activity_id: id,
1912 run_id: None,
1913 completion_token,
1914 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1915 ContentType::Json,
1916 br#""valid""#.to_vec(),
1917 )),
1918 })?;
1919 let result = rx
1920 .recv_timeout(Duration::from_millis(50))
1921 .map_err(|error| ServerError::worker_dispatch("", "", format!("channel: {error}")))?;
1922 assert_eq!(result, Ok(r#""valid""#.to_owned()));
1923 Ok(())
1924 }
1925
1926 #[test]
1927 fn completion_sink_routes_retryable_error() -> Result<(), ServerError> {
1928 let pending = PendingActivities::default();
1929 let workflow_id = WorkflowId::new_v4();
1930 let id = activity_id(3);
1931 let (completion_token, rx) = pending.insert(workflow_id.clone(), id.clone())?;
1932
1933 pending.complete_activity(ActivityCompletion {
1934 workflow_id,
1935 activity_id: id,
1936 run_id: None,
1937 completion_token,
1938 outcome: ActivityCompletionOutcome::Failed(ActivityError {
1939 kind: ActivityErrorKind::Retryable,
1940 message: "temporary".to_owned(),
1941 details: None,
1942 }),
1943 })?;
1944
1945 let result = rx
1946 .recv_timeout(Duration::from_millis(50))
1947 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1948 assert_eq!(result, Err("retryable:temporary".to_owned()));
1949 Ok(())
1950 }
1951
1952 #[test]
1961 fn stale_result_for_other_workflow_does_not_complete_pending_dispatch()
1962 -> Result<(), ServerError> {
1963 let pending = PendingActivities::default();
1964 let post_restart_workflow = WorkflowId::new_v4();
1965 let pre_restart_workflow = WorkflowId::new_v4();
1966 let id = activity_id(1);
1968 let (completion_token, rx) = pending.insert(post_restart_workflow.clone(), id.clone())?;
1969
1970 let rejected = pending.complete_activity(ActivityCompletion {
1972 workflow_id: pre_restart_workflow,
1973 activity_id: id.clone(),
1974 run_id: None,
1975 completion_token: CompletionToken::for_test(),
1976 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1977 ContentType::Json,
1978 br#""stale""#.to_vec(),
1979 )),
1980 });
1981 assert!(matches!(
1982 rejected,
1983 Err(ServerError::ActivityCompletionRejected { .. })
1984 ));
1985 assert!(
1986 rx.try_recv().is_err(),
1987 "stale result for a different workflow must not complete this dispatch"
1988 );
1989
1990 pending.complete_activity(ActivityCompletion {
1992 workflow_id: post_restart_workflow,
1993 activity_id: id,
1994 run_id: None,
1995 completion_token,
1996 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1997 ContentType::Json,
1998 br#""fresh""#.to_vec(),
1999 )),
2000 })?;
2001 let result = rx
2002 .recv_timeout(Duration::from_millis(50))
2003 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
2004 assert_eq!(result, Ok(r#""fresh""#.to_owned()));
2005 Ok(())
2006 }
2007
2008 fn test_tracker() -> HeartbeatTracker {
2011 HeartbeatTracker::new(Duration::from_secs(5))
2012 }
2013
2014 fn greet_request() -> ActivityDispatch {
2017 ActivityDispatch {
2018 namespace: "default".to_owned(),
2019 task_queue: "default".to_owned(),
2020 node: None,
2021 workflow_id: WorkflowId::new_v4(),
2022 run_id: RunId::new_v4(),
2023 activity_id: ActivityId::from_sequence_position(0),
2024 name: "greet".to_owned(),
2025 input: "{}".to_owned(),
2026 config: "{}".to_owned(),
2027 attempt: 1,
2028 labels: std::collections::BTreeMap::new(),
2029 advisory: false,
2030 }
2031 }
2032
2033 #[test]
2034 fn dispatcher_fails_immediately_when_draining_without_workers() {
2035 let registry = ConnectedWorkerRegistry::default();
2036 let drain = DrainState::default();
2037 let dispatcher = WorkerActivityDispatcher::new(registry, "default", test_tracker())
2038 .with_drain_state(drain.clone());
2039
2040 let _ = drain.begin();
2041
2042 let result = dispatcher.dispatch(greet_request());
2043
2044 assert!(result.is_err());
2045 let err = result.err().unwrap_or_default();
2046 assert!(
2047 err.contains("drain"),
2048 "expected drain rejection, got: {err}"
2049 );
2050 }
2051
2052 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2068 async fn dispatch_inside_runtime_task_delivers_promptly_and_round_trips()
2069 -> Result<(), Box<dyn std::error::Error>> {
2070 let registry = ConnectedWorkerRegistry::default();
2071 let pending = PendingActivities::default();
2072 let (worker_tx, mut worker_rx) = tokio::sync::mpsc::channel(32);
2073 let activity_types = [String::from("greet")];
2074 let registration = registry.register("default", activity_types.iter(), worker_tx)?;
2075
2076 let sink = pending.clone();
2077 let echo_worker = tokio::spawn(async move {
2078 let Some(WorkerMessage::ActivityTask(task)) = worker_rx.recv().await else {
2079 return Err("expected an activity task on the worker channel".to_owned());
2080 };
2081 let workflow_id = task
2082 .workflow_id
2083 .ok_or("task missing workflow id")
2084 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
2085 let activity_id = task
2086 .activity_id
2087 .map(ActivityId::from)
2088 .ok_or("task missing activity id")?;
2089 let completion_token =
2090 CompletionToken::from_wire(&workflow_id, &activity_id, task.completion_token)
2091 .map_err(|error| error.to_string())?;
2092 sink.complete_activity(ActivityCompletion {
2093 workflow_id,
2094 activity_id,
2095 run_id: None,
2096 completion_token,
2097 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
2098 ContentType::Json,
2099 br#"{"greeting":"hello"}"#.to_vec(),
2100 )),
2101 })
2102 .map_err(|error| error.to_string())
2103 });
2104
2105 let dispatcher = Arc::new(
2106 WorkerActivityDispatcher::new(registry, "default", test_tracker())
2107 .with_pending(pending),
2108 );
2109 let started = Instant::now();
2110 let dispatch_task = tokio::spawn(futures::future::lazy(move |_| {
2113 dispatcher.dispatch(greet_request())
2114 }));
2115 let result = dispatch_task.await.map_err(|error| error.to_string())?;
2116 let elapsed = started.elapsed();
2117
2118 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
2119 assert!(
2120 elapsed < Duration::from_secs(5),
2121 "dispatch round trip took {elapsed:?}; task delivery must not \
2122 depend on the blocked dispatch thread"
2123 );
2124 echo_worker.await.map_err(|error| error.to_string())??;
2125 registration.deregister()?;
2126 Ok(())
2127 }
2128
2129 #[tokio::test]
2133 async fn dispatch_on_current_thread_runtime_fails_fast()
2134 -> Result<(), Box<dyn std::error::Error>> {
2135 let registry = ConnectedWorkerRegistry::default();
2136 let (worker_tx, _worker_rx) = tokio::sync::mpsc::channel(32);
2137 let activity_types = [String::from("greet")];
2138 let registration = registry.register("default", activity_types.iter(), worker_tx)?;
2139 let dispatcher = WorkerActivityDispatcher::new(registry, "default", test_tracker());
2140
2141 let started = Instant::now();
2142 let result = dispatcher.dispatch(greet_request());
2143 let elapsed = started.elapsed();
2144
2145 let err = result.err().ok_or("expected dispatch to fail")?;
2146 assert!(
2147 err.contains("multi-thread tokio runtime"),
2148 "unexpected error: {err}"
2149 );
2150 assert!(
2151 elapsed < Duration::from_secs(5),
2152 "fail-fast path took {elapsed:?}"
2153 );
2154 registration.deregister()?;
2155 Ok(())
2156 }
2157
2158 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2164 async fn dispatch_pinned_to_node_reaches_only_that_node()
2165 -> Result<(), Box<dyn std::error::Error>> {
2166 let registry = ConnectedWorkerRegistry::default();
2167 let pending = PendingActivities::default();
2168 let activity_types = [String::from("greet")];
2169 let (n1_tx, mut n1_rx) = tokio::sync::mpsc::channel(32);
2170 let (n2_tx, mut n2_rx) = tokio::sync::mpsc::channel(32);
2171 let on_n2 = registry.register_namespaces(
2177 [String::from("default")],
2178 "default",
2179 Some(String::from("n2")),
2180 activity_types.iter(),
2181 n2_tx,
2182 )?;
2183 let on_n1 = registry.register_namespaces(
2184 [String::from("default")],
2185 "default",
2186 Some(String::from("n1")),
2187 activity_types.iter(),
2188 n1_tx,
2189 )?;
2190
2191 let sink = pending.clone();
2195 let echo_n1 = tokio::spawn(async move {
2196 let Some(WorkerMessage::ActivityTask(task)) = n1_rx.recv().await else {
2197 return Err("expected an activity task on the n1 worker channel".to_owned());
2198 };
2199 let workflow_id = task
2200 .workflow_id
2201 .ok_or("task missing workflow id")
2202 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
2203 let activity_id = task
2204 .activity_id
2205 .map(ActivityId::from)
2206 .ok_or("task missing activity id")?;
2207 let completion_token =
2208 CompletionToken::from_wire(&workflow_id, &activity_id, task.completion_token)
2209 .map_err(|error| error.to_string())?;
2210 sink.complete_activity(ActivityCompletion {
2211 workflow_id,
2212 activity_id,
2213 run_id: None,
2214 completion_token,
2215 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
2216 ContentType::Json,
2217 br#"{"greeting":"hello"}"#.to_vec(),
2218 )),
2219 })
2220 .map_err(|error| error.to_string())
2221 });
2222
2223 let dispatcher = Arc::new(
2224 WorkerActivityDispatcher::new(registry.clone(), "default", test_tracker())
2225 .with_pending(pending),
2226 );
2227
2228 let pinned = ActivityDispatch {
2229 node: Some(String::from("n1")),
2230 ..greet_request()
2231 };
2232 let started = Instant::now();
2233 let result = tokio::spawn(futures::future::lazy(move |_| dispatcher.dispatch(pinned)))
2234 .await
2235 .map_err(|error| error.to_string())?;
2236 let elapsed = started.elapsed();
2237
2238 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
2239 assert!(
2240 elapsed < Duration::from_secs(5),
2241 "pinned dispatch round trip took {elapsed:?}; the task must route to n1"
2242 );
2243 echo_n1.await.map_err(|error| error.to_string())??;
2244
2245 assert!(
2247 n2_rx.try_recv().is_err(),
2248 "node=Some(\"n1\") dispatch must not reach the n2 worker"
2249 );
2250
2251 on_n1.deregister()?;
2252 on_n2.deregister()?;
2253 Ok(())
2254 }
2255
2256 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2259 async fn unpinned_dispatch_reaches_a_pooled_worker_regardless_of_node()
2260 -> Result<(), Box<dyn std::error::Error>> {
2261 let registry = ConnectedWorkerRegistry::default();
2262 let pending = PendingActivities::default();
2263 let activity_types = [String::from("greet")];
2264 let (n1_tx, mut n1_rx) = tokio::sync::mpsc::channel(32);
2265 let on_n1 = registry.register_namespaces(
2266 [String::from("default")],
2267 "default",
2268 Some(String::from("n1")),
2269 activity_types.iter(),
2270 n1_tx,
2271 )?;
2272
2273 let sink = pending.clone();
2274 let echo = tokio::spawn(async move {
2275 let Some(WorkerMessage::ActivityTask(task)) = n1_rx.recv().await else {
2276 return Err("expected an activity task on the worker channel".to_owned());
2277 };
2278 let workflow_id = task
2279 .workflow_id
2280 .ok_or("task missing workflow id")
2281 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
2282 let activity_id = task
2283 .activity_id
2284 .map(ActivityId::from)
2285 .ok_or("task missing activity id")?;
2286 let completion_token =
2287 CompletionToken::from_wire(&workflow_id, &activity_id, task.completion_token)
2288 .map_err(|error| error.to_string())?;
2289 sink.complete_activity(ActivityCompletion {
2290 workflow_id,
2291 activity_id,
2292 run_id: None,
2293 completion_token,
2294 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
2295 ContentType::Json,
2296 br#"{"greeting":"hello"}"#.to_vec(),
2297 )),
2298 })
2299 .map_err(|error| error.to_string())
2300 });
2301
2302 let dispatcher = Arc::new(
2303 WorkerActivityDispatcher::new(registry.clone(), "default", test_tracker())
2304 .with_pending(pending),
2305 );
2306
2307 let result = tokio::spawn(futures::future::lazy(move |_| {
2309 dispatcher.dispatch(greet_request())
2310 }))
2311 .await
2312 .map_err(|error| error.to_string())?;
2313
2314 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
2315 echo.await.map_err(|error| error.to_string())??;
2316 on_n1.deregister()?;
2317 Ok(())
2318 }
2319}