1use std::collections::BTreeMap;
60use std::sync::{Arc, OnceLock};
61use std::time::{Duration, Instant};
62
63use aion::{ActivityDispatch, ActivityDispatcher};
64use aion_core::{ActivityId, ContentType, Payload, RunId, WorkflowId};
65use aion_proto::{ProtoActivityId, ProtoActivityTask, ProtoPayload, ProtoWorkflowId};
66use dashmap::DashMap;
67
68use super::dispatch::{ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink};
69use super::heartbeat::{HeartbeatTracker, InFlightActivity};
70use super::registry::{
71 ConnectedWorkerRegistry, WorkerDelivery, WorkerHandle, WorkerId, WorkerMessage,
72};
73use crate::error::ServerError;
74use crate::shutdown::DrainState;
75use tracing::info_span;
76
77type SyncSender = std::sync::mpsc::SyncSender<Result<String, String>>;
78type SyncReceiver = std::sync::mpsc::Receiver<Result<String, String>>;
79
80type PendingActivityKey = (WorkflowId, ActivityId);
98
99pub trait OutboxDeliveryCallback: Send + Sync {
109 fn deliver_completion(
119 &self,
120 workflow_id: &WorkflowId,
121 activity_id: &ActivityId,
122 run_id: Option<&RunId>,
123 result: String,
124 ) -> Result<bool, ServerError>;
125
126 fn deliver_failure(
133 &self,
134 workflow_id: &WorkflowId,
135 activity_id: &ActivityId,
136 run_id: Option<&RunId>,
137 reason: String,
138 ) -> Result<bool, ServerError>;
139}
140
141#[derive(Clone, Default)]
153pub struct PendingActivities {
154 pending: Arc<DashMap<PendingActivityKey, SyncSender>>,
155 outbox_delivery: Arc<OnceLock<Arc<dyn OutboxDeliveryCallback>>>,
156}
157
158impl std::fmt::Debug for PendingActivities {
159 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
160 formatter
161 .debug_struct("PendingActivities")
162 .field("pending", &self.pending.len())
163 .field(
164 "outbox_delivery_installed",
165 &self.outbox_delivery.get().is_some(),
166 )
167 .finish()
168 }
169}
170
171impl PendingActivities {
172 fn insert(&self, workflow_id: WorkflowId, activity_id: ActivityId) -> SyncReceiver {
173 let (tx, rx) = std::sync::mpsc::sync_channel(1);
174 self.pending.insert((workflow_id, activity_id), tx);
175 rx
176 }
177
178 #[cfg(test)]
182 pub(crate) fn insert_for_test(
183 &self,
184 workflow_id: WorkflowId,
185 activity_id: ActivityId,
186 ) -> SyncReceiver {
187 self.insert(workflow_id, activity_id)
188 }
189
190 pub fn set_outbox_delivery(&self, callback: Arc<dyn OutboxDeliveryCallback>) {
196 if self.outbox_delivery.set(callback).is_err() {
197 tracing::warn!("outbox delivery callback already installed; ignoring duplicate set");
198 }
199 }
200
201 fn complete(
209 &self,
210 workflow_id: &WorkflowId,
211 activity_id: &ActivityId,
212 run_id: Option<&RunId>,
213 result: Result<String, String>,
214 ) -> bool {
215 let matched = self
218 .pending
219 .remove(&(workflow_id.clone(), activity_id.clone()));
220 if let Some((_, sender)) = matched {
221 return sender.send(result).is_ok();
222 }
223 let Some(callback) = self.outbox_delivery.get() else {
224 return false;
226 };
227 let outcome = match result {
228 Ok(payload) => callback.deliver_completion(workflow_id, activity_id, run_id, payload),
229 Err(reason) => callback.deliver_failure(workflow_id, activity_id, run_id, reason),
230 };
231 match outcome {
232 Ok(true) => true,
233 Ok(false) => {
234 tracing::debug!(
236 workflow_id = %workflow_id,
237 activity_id = %activity_id,
238 "unmatched outbox completion for a workflow that is not currently live; \
239 recovery will re-arm it"
240 );
241 false
242 }
243 Err(error) => {
244 tracing::warn!(
245 workflow_id = %workflow_id,
246 activity_id = %activity_id,
247 %error,
248 "failed to deliver unmatched outbox completion to the live workflow"
249 );
250 false
251 }
252 }
253 }
254}
255
256impl ActivityCompletionSink for PendingActivities {
257 fn park_activity(
268 &self,
269 workflow_id: &WorkflowId,
270 activity_id: &ActivityId,
271 ) -> Result<(), ServerError> {
272 let matched = self
273 .pending
274 .remove(&(workflow_id.clone(), activity_id.clone()));
275 if let Some((_, sender)) = matched {
276 let _ = sender.send(Err(aion::PARKED_ACTIVITY_REASON.to_owned()));
280 }
281 Ok(())
282 }
283
284 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
285 let result = match completion.outcome {
286 ActivityCompletionOutcome::Succeeded(payload) => {
287 payload_to_string(&payload).map_err(|reason| {
288 tracing::error!(
289 operation = "activity_complete",
290 workflow_id = %completion.workflow_id,
291 activity_id = %completion.activity_id,
292 error_type = "ActivityResultDecode",
293 %reason,
294 "activity completion failed"
295 );
296 ServerError::worker_dispatch("", "", format!("payload decode: {reason}"))
297 })?
298 }
299 ActivityCompletionOutcome::Failed(error) => {
300 let prefix = if error.is_retryable() {
301 "retryable"
302 } else {
303 "terminal"
304 };
305 tracing::error!(
306 operation = "activity_complete",
307 workflow_id = %completion.workflow_id,
308 activity_id = %completion.activity_id,
309 error_type = "ActivityFailed",
310 error_kind = prefix,
311 reason = %error.message,
312 "activity completion failed"
313 );
314 Err(format!("{prefix}:{}", error.message))
315 }
316 };
317 self.complete(
318 &completion.workflow_id,
319 &completion.activity_id,
320 completion.run_id.as_ref(),
321 result,
322 );
323 Ok(())
324 }
325}
326
327fn payload_to_string(payload: &Payload) -> Result<Result<String, String>, String> {
328 match payload.content_type() {
329 ContentType::Json => String::from_utf8(payload.bytes().to_vec())
330 .map(Ok)
331 .map_err(|_| "activity result payload is not valid UTF-8".to_owned()),
332 }
333}
334
335pub struct WorkerActivityDispatcher {
343 registry: ConnectedWorkerRegistry,
344 namespace: String,
345 pending: PendingActivities,
346 heartbeat_tracker: HeartbeatTracker,
347 drain_state: DrainState,
348 tokio_handle: Option<tokio::runtime::Handle>,
349 attempt_owners: Option<super::intervention::AttemptOwnerIndex>,
355}
356
357impl std::fmt::Debug for WorkerActivityDispatcher {
358 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
359 f.debug_struct("WorkerActivityDispatcher")
360 .field("namespace", &self.namespace)
361 .finish_non_exhaustive()
362 }
363}
364
365impl WorkerActivityDispatcher {
366 #[must_use]
374 pub fn new(
375 registry: ConnectedWorkerRegistry,
376 namespace: impl Into<String>,
377 heartbeat_tracker: HeartbeatTracker,
378 ) -> Self {
379 Self {
380 registry,
381 namespace: namespace.into(),
382 pending: PendingActivities::default(),
383 heartbeat_tracker,
384 drain_state: DrainState::default(),
385 tokio_handle: None,
386 attempt_owners: None,
387 }
388 }
389
390 #[must_use]
397 pub fn with_attempt_owners(
398 mut self,
399 attempt_owners: super::intervention::AttemptOwnerIndex,
400 ) -> Self {
401 self.attempt_owners = Some(attempt_owners);
402 self
403 }
404
405 #[must_use]
407 pub fn with_pending(mut self, pending: PendingActivities) -> Self {
408 self.pending = pending;
409 self
410 }
411
412 #[must_use]
414 pub fn with_drain_state(mut self, drain_state: DrainState) -> Self {
415 self.drain_state = drain_state;
416 self
417 }
418
419 #[must_use]
421 pub fn with_tokio_handle(mut self, tokio_handle: tokio::runtime::Handle) -> Self {
422 self.tokio_handle = Some(tokio_handle);
423 self
424 }
425}
426
427impl WorkerActivityDispatcher {
428 fn ensure_accepting(
429 &self,
430 namespace: &str,
431 activity_type: &str,
432 workflow_id: &WorkflowId,
433 activity_id: &ActivityId,
434 worker_id: Option<WorkerId>,
435 ) -> Result<(), String> {
436 self.drain_state
437 .ensure_accepting(namespace, activity_type)
438 .map_err(|error| {
439 let reason = error.to_string();
440 log_worker_error(
441 "WorkerDispatch",
442 namespace,
443 activity_type,
444 workflow_id,
445 activity_id,
446 worker_id,
447 &reason,
448 );
449 reason
450 })
451 }
452
453 fn select_worker_or_wait(
457 &self,
458 namespace: &str,
459 task_queue: &str,
460 activity_type: &str,
461 node: Option<&str>,
462 workflow_id: &WorkflowId,
463 activity_id: &ActivityId,
464 ) -> Result<WorkerHandle, String> {
465 loop {
466 match self
473 .registry
474 .select_worker(namespace, task_queue, activity_type, node)
475 {
476 Ok(Some(worker)) => return Ok(worker),
477 Ok(None) => {
478 self.ensure_accepting(
479 namespace,
480 activity_type,
481 workflow_id,
482 activity_id,
483 None,
484 )?;
485 tracing::info!(
486 namespace,
487 activity_type,
488 node,
489 workflow_id = %workflow_id,
490 activity_id = %activity_id,
491 "no connected worker; waiting for a matching worker to register"
492 );
493 match &self.tokio_handle {
494 Some(handle) => {
495 handle.block_on(self.registry.wait_for_worker());
496 }
497 None => match tokio::runtime::Handle::try_current() {
498 Ok(handle) => {
499 handle.block_on(self.registry.wait_for_worker());
500 }
501 Err(_) => {
502 std::thread::sleep(Duration::from_millis(500));
503 }
504 },
505 }
506 }
507 Err(error) => {
508 let reason = format!("registry error: {error}");
509 log_worker_error(
510 "WorkerRegistry",
511 namespace,
512 activity_type,
513 workflow_id,
514 activity_id,
515 None,
516 &reason,
517 );
518 return Err(reason);
519 }
520 }
521 }
522 }
523
524 fn track_worker_task(
525 &self,
526 worker_id: WorkerId,
527 activity_type: &str,
528 workflow_id: &WorkflowId,
529 activity_id: &ActivityId,
530 ) -> Result<(), String> {
531 self.heartbeat_tracker
532 .track_task(
533 worker_id,
534 InFlightActivity {
535 workflow_id: workflow_id.clone(),
536 activity_id: activity_id.clone(),
537 },
538 Instant::now(),
539 )
540 .map_err(|error| {
541 let reason = error.to_string();
542 log_worker_error(
543 "WorkerHeartbeatTracker",
544 &self.namespace,
545 activity_type,
546 workflow_id,
547 activity_id,
548 Some(worker_id),
549 &reason,
550 );
551 reason
552 })
553 }
554
555 fn cleanup_activity(
556 &self,
557 worker_id: WorkerId,
558 workflow_id: &WorkflowId,
559 activity_id: &ActivityId,
560 ) {
561 self.pending
562 .pending
563 .remove(&(workflow_id.clone(), activity_id.clone()));
564 let _ = self
565 .heartbeat_tracker
566 .complete_task(worker_id, workflow_id, activity_id);
567 self.drain_state.notify_activity_drained();
568 }
569
570 fn send_activity_task(
580 &self,
581 worker: &WorkerHandle,
582 task: ProtoActivityTask,
583 activity_type: &str,
584 workflow_id: &WorkflowId,
585 activity_id: &ActivityId,
586 ) -> Result<(), String> {
587 match worker.delivery() {
588 WorkerDelivery::Grpc(sender) => {
589 match sender.try_send(WorkerMessage::ActivityTask(task)) {
590 Ok(()) => Ok(()),
591 Err(error) => {
592 let worker_id = worker.id();
593 let reason = format!("worker task channel full or closed: {error}");
594 self.cleanup_activity(worker_id, workflow_id, activity_id);
595 log_worker_error(
596 "WorkerChannelClosed",
597 &self.namespace,
598 activity_type,
599 workflow_id,
600 activity_id,
601 Some(worker_id),
602 &reason,
603 );
604 Err(reason)
605 }
606 }
607 }
608 #[cfg(feature = "liminal-transport")]
609 WorkerDelivery::Liminal(delivery) => self.send_liminal_activity_task(
610 worker.id(),
611 delivery,
612 task,
613 activity_type,
614 workflow_id,
615 activity_id,
616 ),
617 }
618 }
619
620 #[cfg(feature = "liminal-transport")]
644 fn send_liminal_activity_task(
645 &self,
646 worker_id: WorkerId,
647 delivery: &super::liminal_transport::LiminalWorkerDelivery,
648 task: ProtoActivityTask,
649 activity_type: &str,
650 workflow_id: &WorkflowId,
651 activity_id: &ActivityId,
652 ) -> Result<(), String> {
653 let heartbeat_window_ms =
654 u64::try_from(self.heartbeat_tracker.heartbeat_window().as_millis())
655 .unwrap_or(u64::MAX);
656 let attempt = task.attempt;
657 let request = super::liminal_transport::DispatchRequest {
658 activity_type: activity_type.to_owned(),
659 workflow_id: workflow_id.clone(),
660 ordinal: activity_id.sequence_position(),
661 run_id: None,
662 attempt,
663 labels: task.labels.into_iter().collect(),
664 heartbeat_window_ms,
665 input: task.input.map(|payload| payload.bytes).unwrap_or_default(),
666 };
667 let awaiter = match delivery.push_dispatch(&request) {
671 Ok(awaiter) => awaiter,
672 Err(error) => {
673 let reason = format!("worker liminal push failed: {error}");
674 self.cleanup_activity(worker_id, workflow_id, activity_id);
675 log_worker_error(
676 "WorkerChannelClosed",
677 &self.namespace,
678 activity_type,
679 workflow_id,
680 activity_id,
681 Some(worker_id),
682 &reason,
683 );
684 return Err(reason);
685 }
686 };
687 let owner_binding = self.attempt_owners.as_ref().map(|owners| {
691 super::liminal_transport::AttemptOwnerGuard::bind(
692 owners.clone(),
693 super::intervention::AttemptKey::new(
694 workflow_id.clone(),
695 activity_id.clone(),
696 attempt,
697 ),
698 worker_id,
699 )
700 });
701 self.spawn_liminal_reply_router(
702 worker_id,
703 awaiter,
704 workflow_id,
705 activity_id,
706 owner_binding,
707 );
708 Ok(())
709 }
710
711 #[cfg(feature = "liminal-transport")]
751 fn spawn_liminal_reply_router(
752 &self,
753 worker_id: WorkerId,
754 awaiter: liminal_server::server::connection::PushReplyAwaiter,
755 workflow_id: &WorkflowId,
756 activity_id: &ActivityId,
757 owner_binding: Option<super::liminal_transport::AttemptOwnerGuard>,
758 ) {
759 let pending = self.pending.clone();
760 let heartbeat_tracker = self.heartbeat_tracker.clone();
761 let drain_state = self.drain_state.clone();
762 let workflow_id = workflow_id.clone();
763 let activity_id = activity_id.clone();
764 std::thread::spawn(move || {
765 let _owner_binding = owner_binding;
769 route_liminal_reply(
770 &pending,
771 &heartbeat_tracker,
772 &drain_state,
773 worker_id,
774 &awaiter,
775 &workflow_id,
776 &activity_id,
777 );
778 });
779 }
780
781 fn await_activity_result(
785 &self,
786 context: &ActivityDispatchContext<'_>,
787 rx: &SyncReceiver,
788 ) -> Result<String, String> {
789 match self.registry.is_registered(context.worker_id) {
797 Ok(true) => {}
798 Ok(false) => {
799 if let Ok(result) = rx.try_recv() {
803 return self.deliver_result(context, result);
804 }
805 self.cleanup_activity(context.worker_id, context.workflow_id, context.activity_id);
806 let reason = format!(
807 "retryable:{}",
808 super::dispatch::lost_worker_error(context.worker_id).message
809 );
810 log_worker_error(
811 "WorkerLost",
812 &self.namespace,
813 context.activity_type,
814 context.workflow_id,
815 context.activity_id,
816 Some(context.worker_id),
817 &reason,
818 );
819 return Err(reason);
820 }
821 Err(error) => {
822 self.cleanup_activity(context.worker_id, context.workflow_id, context.activity_id);
823 let reason = format!("worker registry inspection failed: {error}");
824 log_worker_error(
825 "WorkerRegistry",
826 &self.namespace,
827 context.activity_type,
828 context.workflow_id,
829 context.activity_id,
830 Some(context.worker_id),
831 &reason,
832 );
833 return Err(reason);
834 }
835 }
836 if let Ok(result) = rx.recv() {
837 return self.deliver_result(context, result);
838 }
839 self.cleanup_activity(context.worker_id, context.workflow_id, context.activity_id);
842 let reason = "activity response channel dropped".to_owned();
843 log_worker_error(
844 "WorkerChannelClosed",
845 &self.namespace,
846 context.activity_type,
847 context.workflow_id,
848 context.activity_id,
849 Some(context.worker_id),
850 &reason,
851 );
852 Err(reason)
853 }
854
855 fn deliver_result(
856 &self,
857 context: &ActivityDispatchContext<'_>,
858 result: Result<String, String>,
859 ) -> Result<String, String> {
860 self.pending
861 .pending
862 .remove(&(context.workflow_id.clone(), context.activity_id.clone()));
863 if let Err(reason) = &result
868 && aion::is_parked_reason(reason)
869 {
870 tracing::info!(
871 operation = "activity_dispatch",
872 namespace = %self.namespace,
873 workflow_id = %context.workflow_id,
874 activity_id = %context.activity_id,
875 activity_type = context.activity_type,
876 worker_id = ?context.worker_id,
877 "activity parked for restart recovery"
878 );
879 return result;
880 }
881 log_activity_completion(context, result.is_ok());
882 result.inspect_err(|reason| {
883 log_worker_error(
884 "ActivityFailed",
885 &self.namespace,
886 context.activity_type,
887 context.workflow_id,
888 context.activity_id,
889 Some(context.worker_id),
890 reason,
891 );
892 })
893 }
894}
895
896impl ActivityDispatcher for WorkerActivityDispatcher {
897 fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
898 match tokio::runtime::Handle::try_current() {
899 Ok(handle) => match handle.runtime_flavor() {
900 tokio::runtime::RuntimeFlavor::MultiThread => {
901 tokio::task::block_in_place(|| self.dispatch_blocking(request))
908 }
909 flavor => Err(format!(
910 "activity dispatch blocks the calling thread until the worker responds; \
911 a {flavor:?} tokio runtime cannot host that wait because the worker \
912 stream forwarder shares its only executor thread and the task could \
913 never be delivered — run the engine on a multi-thread tokio runtime"
914 )),
915 },
916 Err(_) => self.dispatch_blocking(request),
920 }
921 }
922}
923
924impl WorkerActivityDispatcher {
925 fn dispatch_blocking(&self, request: ActivityDispatch) -> Result<String, String> {
942 let ActivityDispatch {
943 namespace,
944 task_queue,
945 node,
949 workflow_id,
950 activity_id,
951 name,
952 input,
953 config: _,
954 attempt,
955 labels,
956 } = request;
957 let started_at = Instant::now();
958 self.ensure_accepting(&namespace, &name, &workflow_id, &activity_id, None)?;
959 let worker = self.select_worker_or_wait(
960 &namespace,
961 &task_queue,
962 &name,
963 node.as_deref(),
964 &workflow_id,
965 &activity_id,
966 )?;
967 let worker_id = worker.id();
968 let span = info_span!(
969 "activity_dispatch",
970 operation = "activity_dispatch",
971 namespace = %namespace,
972 task_queue = %task_queue,
973 node = node.as_deref(),
974 workflow_id = %workflow_id,
975 activity_id = %activity_id,
976 activity_type = %name,
977 worker_id = ?worker_id,
978 );
979 let _span_guard = span.enter();
980 self.ensure_accepting(
981 &namespace,
982 &name,
983 &workflow_id,
984 &activity_id,
985 Some(worker_id),
986 )?;
987
988 let task = activity_task(&name, &input, &workflow_id, &activity_id, attempt, labels);
989 let rx = self
990 .pending
991 .insert(workflow_id.clone(), activity_id.clone());
992 self.track_worker_task(worker_id, &name, &workflow_id, &activity_id)?;
993 self.send_activity_task(&worker, task, &name, &workflow_id, &activity_id)?;
994 let context = ActivityDispatchContext {
995 namespace: &namespace,
996 activity_type: &name,
997 worker_id,
998 workflow_id: &workflow_id,
999 activity_id: &activity_id,
1000 started_at,
1001 };
1002 self.await_activity_result(&context, &rx)
1003 }
1004}
1005
1006#[cfg(feature = "liminal-transport")]
1013fn route_liminal_reply(
1014 pending: &PendingActivities,
1015 heartbeat_tracker: &HeartbeatTracker,
1016 drain_state: &DrainState,
1017 worker_id: WorkerId,
1018 awaiter: &liminal_server::server::connection::PushReplyAwaiter,
1019 workflow_id: &WorkflowId,
1020 activity_id: &ActivityId,
1021) {
1022 let waited = super::liminal_transport::receive_bridge_reply(awaiter, || {
1027 heartbeat_tracker
1028 .is_tracked(worker_id, workflow_id, activity_id)
1029 .unwrap_or(false)
1030 });
1031 let (run_id, outcome, synthesized) = match waited {
1035 Ok(Some(response)) => (response.run_id, response.outcome, false),
1036 Ok(None) => {
1037 tracing::debug!(
1038 worker_id = ?worker_id,
1039 workflow_id = %workflow_id,
1040 activity_id = %activity_id,
1041 "liminal dispatch resolved by another path; abandoning reply wait"
1042 );
1043 return;
1044 }
1045 Err(error) if error.is_worker_connection_lost() => (
1046 None,
1047 Err(format!(
1048 "retryable:{}",
1049 super::dispatch::lost_worker_error(worker_id).message
1050 )),
1051 true,
1052 ),
1053 Err(error) => (
1054 None,
1055 Err(format!("retryable:worker liminal reply failed: {error}")),
1056 true,
1057 ),
1058 };
1059 let was_tracked = heartbeat_tracker
1064 .complete_task(worker_id, workflow_id, activity_id)
1065 .unwrap_or_else(|error| {
1066 tracing::error!(
1067 worker_id = ?worker_id,
1068 workflow_id = %workflow_id,
1069 activity_id = %activity_id,
1070 %error,
1071 "failed to clear in-flight tracking for completed liminal activity"
1072 );
1073 true
1074 });
1075 if synthesized && !was_tracked {
1076 tracing::debug!(
1081 worker_id = ?worker_id,
1082 workflow_id = %workflow_id,
1083 activity_id = %activity_id,
1084 "liminal dispatch already resolved; dropping synthesized lost-worker failure"
1085 );
1086 return;
1087 }
1088 drain_state.notify_activity_drained();
1089 pending.complete(workflow_id, activity_id, run_id.as_ref(), outcome);
1090}
1091
1092struct ActivityDispatchContext<'a> {
1093 namespace: &'a str,
1094 activity_type: &'a str,
1095 worker_id: WorkerId,
1096 workflow_id: &'a WorkflowId,
1097 activity_id: &'a ActivityId,
1098 started_at: Instant,
1099}
1100
1101fn activity_task(
1102 activity_type: &str,
1103 input: &str,
1104 workflow_id: &WorkflowId,
1105 activity_id: &ActivityId,
1106 attempt: u32,
1107 labels: BTreeMap<String, String>,
1108) -> ProtoActivityTask {
1109 ProtoActivityTask {
1110 workflow_id: Some(ProtoWorkflowId::from(workflow_id.clone())),
1111 activity_id: Some(ProtoActivityId::from(activity_id.clone())),
1112 activity_type: activity_type.to_owned(),
1113 input: Some(ProtoPayload {
1114 content_type: String::from("application/json"),
1115 bytes: input.as_bytes().to_vec(),
1116 }),
1117 attempt,
1118 labels: labels.into_iter().collect(),
1119 run_id: None,
1122 }
1123}
1124
1125fn log_activity_completion(context: &ActivityDispatchContext<'_>, succeeded: bool) {
1126 let duration_ms = duration_ms(context.started_at.elapsed());
1127 tracing::info!(
1128 operation = "activity_complete",
1129 namespace = context.namespace,
1130 workflow_id = %context.workflow_id,
1131 activity_id = %context.activity_id,
1132 activity_type = context.activity_type,
1133 worker_id = ?context.worker_id,
1134 duration_ms,
1135 outcome = if succeeded { "succeeded" } else { "failed" },
1136 "activity completed"
1137 );
1138}
1139
1140fn duration_ms(duration: Duration) -> u64 {
1141 u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
1142}
1143
1144fn log_worker_error(
1145 error_type: &'static str,
1146 namespace: &str,
1147 activity_type: &str,
1148 workflow_id: &WorkflowId,
1149 activity_id: &ActivityId,
1150 worker_id: Option<super::registry::WorkerId>,
1151 reason: &str,
1152) {
1153 tracing::error!(
1154 operation = "activity_dispatch",
1155 namespace,
1156 workflow_id = %workflow_id,
1157 activity_id = %activity_id,
1158 activity_type,
1159 worker_id = ?worker_id,
1160 error_type,
1161 reason,
1162 "worker interaction failed"
1163 );
1164}
1165
1166#[cfg(test)]
1167mod tests {
1168 use std::sync::Mutex;
1169
1170 use aion_core::{ActivityError, ActivityErrorKind, ContentType, Payload};
1171
1172 use super::*;
1173
1174 fn activity_id(pos: u64) -> ActivityId {
1175 ActivityId::from_sequence_position(pos)
1176 }
1177
1178 #[test]
1179 fn pending_insert_and_complete_delivers_result() {
1180 let pending = PendingActivities::default();
1181 let workflow_id = WorkflowId::new_v4();
1182 let id = activity_id(1);
1183 let rx = pending.insert(workflow_id.clone(), id.clone());
1184
1185 assert!(pending.complete(&workflow_id, &id, None, Ok("done".to_owned())));
1186 assert_eq!(
1187 rx.recv_timeout(Duration::from_millis(50)),
1188 Ok(Ok("done".to_owned()))
1189 );
1190 }
1191
1192 #[test]
1193 fn pending_complete_unknown_returns_false() {
1194 let pending = PendingActivities::default();
1195 assert!(!pending.complete(
1196 &WorkflowId::new_v4(),
1197 &activity_id(99),
1198 None,
1199 Ok("orphan".to_owned())
1200 ));
1201 }
1202
1203 #[derive(Default)]
1204 struct RecordingOutboxCallback {
1205 completions: Mutex<Vec<(WorkflowId, ActivityId, String)>>,
1206 failures: Mutex<Vec<(WorkflowId, ActivityId, String)>>,
1207 live: bool,
1208 }
1209
1210 impl OutboxDeliveryCallback for RecordingOutboxCallback {
1211 fn deliver_completion(
1212 &self,
1213 workflow_id: &WorkflowId,
1214 activity_id: &ActivityId,
1215 run_id: Option<&RunId>,
1216 result: String,
1217 ) -> Result<bool, ServerError> {
1218 let _ = run_id;
1219 self.completions
1220 .lock()
1221 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1222 .push((workflow_id.clone(), activity_id.clone(), result));
1223 Ok(self.live)
1224 }
1225
1226 fn deliver_failure(
1227 &self,
1228 workflow_id: &WorkflowId,
1229 activity_id: &ActivityId,
1230 run_id: Option<&RunId>,
1231 reason: String,
1232 ) -> Result<bool, ServerError> {
1233 let _ = run_id;
1234 self.failures
1235 .lock()
1236 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1237 .push((workflow_id.clone(), activity_id.clone(), reason));
1238 Ok(self.live)
1239 }
1240 }
1241
1242 #[test]
1243 fn unmatched_completion_routes_to_outbox_callback_when_installed() -> Result<(), ServerError> {
1244 let pending = PendingActivities::default();
1245 let callback = Arc::new(RecordingOutboxCallback {
1246 live: true,
1247 ..RecordingOutboxCallback::default()
1248 });
1249 pending.clone().set_outbox_delivery(callback.clone());
1251
1252 let workflow_id = WorkflowId::new_v4();
1253 let id = activity_id(7);
1254
1255 assert!(pending.complete(&workflow_id, &id, None, Ok("done".to_owned())));
1258 let completions = callback
1259 .completions
1260 .lock()
1261 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?;
1262 assert_eq!(completions.len(), 1);
1263 assert_eq!(completions[0].0, workflow_id);
1264 assert_eq!(completions[0].1, id);
1265 assert_eq!(completions[0].2, "done");
1266 Ok(())
1267 }
1268
1269 #[test]
1270 fn unmatched_failure_routes_to_outbox_callback_and_not_live_reports_false()
1271 -> Result<(), ServerError> {
1272 let pending = PendingActivities::default();
1273 let callback = Arc::new(RecordingOutboxCallback::default());
1275 pending.set_outbox_delivery(callback.clone());
1276
1277 let workflow_id = WorkflowId::new_v4();
1278 let id = activity_id(8);
1279
1280 assert!(!pending.complete(&workflow_id, &id, None, Err("retryable:boom".to_owned())));
1281 let failures = callback
1282 .failures
1283 .lock()
1284 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?;
1285 assert_eq!(failures.len(), 1);
1286 assert_eq!(failures[0].2, "retryable:boom");
1287 Ok(())
1288 }
1289
1290 #[test]
1291 fn unmatched_completion_is_silent_drop_when_no_callback_installed() {
1292 let pending = PendingActivities::default();
1295 assert!(!pending.complete(
1296 &WorkflowId::new_v4(),
1297 &activity_id(9),
1298 None,
1299 Ok("x".to_owned())
1300 ));
1301 }
1302
1303 #[test]
1304 fn matched_completion_never_reaches_outbox_callback() -> Result<(), ServerError> {
1305 let pending = PendingActivities::default();
1306 let callback = Arc::new(RecordingOutboxCallback {
1307 live: true,
1308 ..RecordingOutboxCallback::default()
1309 });
1310 pending.set_outbox_delivery(callback.clone());
1311
1312 let workflow_id = WorkflowId::new_v4();
1313 let id = activity_id(10);
1314 let rx = pending.insert(workflow_id.clone(), id.clone());
1315
1316 assert!(pending.complete(&workflow_id, &id, None, Ok("matched".to_owned())));
1317 assert_eq!(
1318 rx.recv_timeout(Duration::from_millis(50)),
1319 Ok(Ok("matched".to_owned()))
1320 );
1321 assert!(
1322 callback
1323 .completions
1324 .lock()
1325 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1326 .is_empty(),
1327 "a matched completion must deliver to its waiter, not the outbox callback"
1328 );
1329 Ok(())
1330 }
1331
1332 #[test]
1336 fn park_activity_resolves_matched_waiter_with_the_parked_sentinel() -> Result<(), ServerError> {
1337 let pending = PendingActivities::default();
1338 let workflow_id = WorkflowId::new_v4();
1339 let id = activity_id(11);
1340 let rx = pending.insert(workflow_id.clone(), id.clone());
1341
1342 pending.park_activity(&workflow_id, &id)?;
1343 let result = rx
1344 .recv_timeout(Duration::from_millis(50))
1345 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1346 assert_eq!(result, Err(aion::PARKED_ACTIVITY_REASON.to_owned()));
1347 Ok(())
1348 }
1349
1350 #[test]
1354 fn unmatched_park_is_a_noop_and_never_reaches_the_outbox_callback() -> Result<(), ServerError> {
1355 let pending = PendingActivities::default();
1356 let callback = Arc::new(RecordingOutboxCallback {
1357 live: true,
1358 ..RecordingOutboxCallback::default()
1359 });
1360 pending.set_outbox_delivery(callback.clone());
1361
1362 pending.park_activity(&WorkflowId::new_v4(), &activity_id(12))?;
1363
1364 assert!(
1365 callback
1366 .failures
1367 .lock()
1368 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1369 .is_empty(),
1370 "a park must never be delivered as an outbox failure"
1371 );
1372 assert!(
1373 callback
1374 .completions
1375 .lock()
1376 .map_err(|_| ServerError::lock_poisoned("recording outbox callback"))?
1377 .is_empty(),
1378 "a park must never be delivered as an outbox completion"
1379 );
1380 Ok(())
1381 }
1382
1383 #[test]
1384 fn completion_sink_routes_success() -> Result<(), ServerError> {
1385 let pending = PendingActivities::default();
1386 let workflow_id = WorkflowId::new_v4();
1387 let id = activity_id(2);
1388 let rx = pending.insert(workflow_id.clone(), id.clone());
1389 let payload = Payload::new(ContentType::Json, br#"{"greeting":"hi"}"#.to_vec());
1390
1391 pending.complete_activity(ActivityCompletion {
1392 workflow_id,
1393 activity_id: id,
1394 run_id: None,
1395 outcome: ActivityCompletionOutcome::Succeeded(payload),
1396 })?;
1397
1398 let result = rx
1399 .recv_timeout(Duration::from_millis(50))
1400 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1401 assert_eq!(result, Ok(r#"{"greeting":"hi"}"#.to_owned()));
1402 Ok(())
1403 }
1404
1405 #[test]
1406 fn completion_sink_routes_retryable_error() -> Result<(), ServerError> {
1407 let pending = PendingActivities::default();
1408 let workflow_id = WorkflowId::new_v4();
1409 let id = activity_id(3);
1410 let rx = pending.insert(workflow_id.clone(), id.clone());
1411
1412 pending.complete_activity(ActivityCompletion {
1413 workflow_id,
1414 activity_id: id,
1415 run_id: None,
1416 outcome: ActivityCompletionOutcome::Failed(ActivityError {
1417 kind: ActivityErrorKind::Retryable,
1418 message: "temporary".to_owned(),
1419 details: None,
1420 }),
1421 })?;
1422
1423 let result = rx
1424 .recv_timeout(Duration::from_millis(50))
1425 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1426 assert_eq!(result, Err("retryable:temporary".to_owned()));
1427 Ok(())
1428 }
1429
1430 #[test]
1439 fn stale_result_for_other_workflow_does_not_complete_pending_dispatch()
1440 -> Result<(), ServerError> {
1441 let pending = PendingActivities::default();
1442 let post_restart_workflow = WorkflowId::new_v4();
1443 let pre_restart_workflow = WorkflowId::new_v4();
1444 let id = activity_id(1);
1446 let rx = pending.insert(post_restart_workflow.clone(), id.clone());
1447
1448 pending.complete_activity(ActivityCompletion {
1450 workflow_id: pre_restart_workflow,
1451 activity_id: id.clone(),
1452 run_id: None,
1453 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1454 ContentType::Json,
1455 br#""stale""#.to_vec(),
1456 )),
1457 })?;
1458 assert!(
1459 rx.try_recv().is_err(),
1460 "stale result for a different workflow must not complete this dispatch"
1461 );
1462
1463 pending.complete_activity(ActivityCompletion {
1465 workflow_id: post_restart_workflow,
1466 activity_id: id,
1467 run_id: None,
1468 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1469 ContentType::Json,
1470 br#""fresh""#.to_vec(),
1471 )),
1472 })?;
1473 let result = rx
1474 .recv_timeout(Duration::from_millis(50))
1475 .map_err(|e| ServerError::worker_dispatch("", "", format!("channel: {e}")))?;
1476 assert_eq!(result, Ok(r#""fresh""#.to_owned()));
1477 Ok(())
1478 }
1479
1480 fn test_tracker() -> HeartbeatTracker {
1483 HeartbeatTracker::new(Duration::from_secs(5))
1484 }
1485
1486 fn greet_request() -> ActivityDispatch {
1489 ActivityDispatch {
1490 namespace: "default".to_owned(),
1491 task_queue: "default".to_owned(),
1492 node: None,
1493 workflow_id: WorkflowId::new_v4(),
1494 activity_id: ActivityId::from_sequence_position(0),
1495 name: "greet".to_owned(),
1496 input: "{}".to_owned(),
1497 config: "{}".to_owned(),
1498 attempt: 1,
1499 labels: std::collections::BTreeMap::new(),
1500 }
1501 }
1502
1503 #[test]
1504 fn dispatcher_fails_immediately_when_draining_without_workers() {
1505 let registry = ConnectedWorkerRegistry::default();
1506 let drain = DrainState::default();
1507 let dispatcher = WorkerActivityDispatcher::new(registry, "default", test_tracker())
1508 .with_drain_state(drain.clone());
1509
1510 let _ = drain.begin();
1511
1512 let result = dispatcher.dispatch(greet_request());
1513
1514 assert!(result.is_err());
1515 let err = result.err().unwrap_or_default();
1516 assert!(
1517 err.contains("drain"),
1518 "expected drain rejection, got: {err}"
1519 );
1520 }
1521
1522 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1538 async fn dispatch_inside_runtime_task_delivers_promptly_and_round_trips()
1539 -> Result<(), Box<dyn std::error::Error>> {
1540 let registry = ConnectedWorkerRegistry::default();
1541 let pending = PendingActivities::default();
1542 let (worker_tx, mut worker_rx) = tokio::sync::mpsc::channel(32);
1543 let activity_types = [String::from("greet")];
1544 let registration = registry.register("default", activity_types.iter(), worker_tx)?;
1545
1546 let sink = pending.clone();
1547 let echo_worker = tokio::spawn(async move {
1548 let Some(WorkerMessage::ActivityTask(task)) = worker_rx.recv().await else {
1549 return Err("expected an activity task on the worker channel".to_owned());
1550 };
1551 let workflow_id = task
1552 .workflow_id
1553 .ok_or("task missing workflow id")
1554 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
1555 let activity_id = task
1556 .activity_id
1557 .map(ActivityId::from)
1558 .ok_or("task missing activity id")?;
1559 sink.complete_activity(ActivityCompletion {
1560 workflow_id,
1561 activity_id,
1562 run_id: None,
1563 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1564 ContentType::Json,
1565 br#"{"greeting":"hello"}"#.to_vec(),
1566 )),
1567 })
1568 .map_err(|error| error.to_string())
1569 });
1570
1571 let dispatcher = Arc::new(
1572 WorkerActivityDispatcher::new(registry, "default", test_tracker())
1573 .with_pending(pending),
1574 );
1575 let started = Instant::now();
1576 let dispatch_task = tokio::spawn(futures::future::lazy(move |_| {
1579 dispatcher.dispatch(greet_request())
1580 }));
1581 let result = dispatch_task.await.map_err(|error| error.to_string())?;
1582 let elapsed = started.elapsed();
1583
1584 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
1585 assert!(
1586 elapsed < Duration::from_secs(5),
1587 "dispatch round trip took {elapsed:?}; task delivery must not \
1588 depend on the blocked dispatch thread"
1589 );
1590 echo_worker.await.map_err(|error| error.to_string())??;
1591 registration.deregister()?;
1592 Ok(())
1593 }
1594
1595 #[tokio::test]
1599 async fn dispatch_on_current_thread_runtime_fails_fast()
1600 -> Result<(), Box<dyn std::error::Error>> {
1601 let registry = ConnectedWorkerRegistry::default();
1602 let (worker_tx, _worker_rx) = tokio::sync::mpsc::channel(32);
1603 let activity_types = [String::from("greet")];
1604 let registration = registry.register("default", activity_types.iter(), worker_tx)?;
1605 let dispatcher = WorkerActivityDispatcher::new(registry, "default", test_tracker());
1606
1607 let started = Instant::now();
1608 let result = dispatcher.dispatch(greet_request());
1609 let elapsed = started.elapsed();
1610
1611 let err = result.err().ok_or("expected dispatch to fail")?;
1612 assert!(
1613 err.contains("multi-thread tokio runtime"),
1614 "unexpected error: {err}"
1615 );
1616 assert!(
1617 elapsed < Duration::from_secs(5),
1618 "fail-fast path took {elapsed:?}"
1619 );
1620 registration.deregister()?;
1621 Ok(())
1622 }
1623
1624 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1630 async fn dispatch_pinned_to_node_reaches_only_that_node()
1631 -> Result<(), Box<dyn std::error::Error>> {
1632 let registry = ConnectedWorkerRegistry::default();
1633 let pending = PendingActivities::default();
1634 let activity_types = [String::from("greet")];
1635 let (n1_tx, mut n1_rx) = tokio::sync::mpsc::channel(32);
1636 let (n2_tx, mut n2_rx) = tokio::sync::mpsc::channel(32);
1637 let on_n2 = registry.register_namespaces(
1643 [String::from("default")],
1644 "default",
1645 Some(String::from("n2")),
1646 activity_types.iter(),
1647 n2_tx,
1648 )?;
1649 let on_n1 = registry.register_namespaces(
1650 [String::from("default")],
1651 "default",
1652 Some(String::from("n1")),
1653 activity_types.iter(),
1654 n1_tx,
1655 )?;
1656
1657 let sink = pending.clone();
1661 let echo_n1 = tokio::spawn(async move {
1662 let Some(WorkerMessage::ActivityTask(task)) = n1_rx.recv().await else {
1663 return Err("expected an activity task on the n1 worker channel".to_owned());
1664 };
1665 let workflow_id = task
1666 .workflow_id
1667 .ok_or("task missing workflow id")
1668 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
1669 let activity_id = task
1670 .activity_id
1671 .map(ActivityId::from)
1672 .ok_or("task missing activity id")?;
1673 sink.complete_activity(ActivityCompletion {
1674 workflow_id,
1675 activity_id,
1676 run_id: None,
1677 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1678 ContentType::Json,
1679 br#"{"greeting":"hello"}"#.to_vec(),
1680 )),
1681 })
1682 .map_err(|error| error.to_string())
1683 });
1684
1685 let dispatcher = Arc::new(
1686 WorkerActivityDispatcher::new(registry.clone(), "default", test_tracker())
1687 .with_pending(pending),
1688 );
1689
1690 let pinned = ActivityDispatch {
1691 node: Some(String::from("n1")),
1692 ..greet_request()
1693 };
1694 let started = Instant::now();
1695 let result = tokio::spawn(futures::future::lazy(move |_| dispatcher.dispatch(pinned)))
1696 .await
1697 .map_err(|error| error.to_string())?;
1698 let elapsed = started.elapsed();
1699
1700 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
1701 assert!(
1702 elapsed < Duration::from_secs(5),
1703 "pinned dispatch round trip took {elapsed:?}; the task must route to n1"
1704 );
1705 echo_n1.await.map_err(|error| error.to_string())??;
1706
1707 assert!(
1709 n2_rx.try_recv().is_err(),
1710 "node=Some(\"n1\") dispatch must not reach the n2 worker"
1711 );
1712
1713 on_n1.deregister()?;
1714 on_n2.deregister()?;
1715 Ok(())
1716 }
1717
1718 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1721 async fn unpinned_dispatch_reaches_a_pooled_worker_regardless_of_node()
1722 -> Result<(), Box<dyn std::error::Error>> {
1723 let registry = ConnectedWorkerRegistry::default();
1724 let pending = PendingActivities::default();
1725 let activity_types = [String::from("greet")];
1726 let (n1_tx, mut n1_rx) = tokio::sync::mpsc::channel(32);
1727 let on_n1 = registry.register_namespaces(
1728 [String::from("default")],
1729 "default",
1730 Some(String::from("n1")),
1731 activity_types.iter(),
1732 n1_tx,
1733 )?;
1734
1735 let sink = pending.clone();
1736 let echo = tokio::spawn(async move {
1737 let Some(WorkerMessage::ActivityTask(task)) = n1_rx.recv().await else {
1738 return Err("expected an activity task on the worker channel".to_owned());
1739 };
1740 let workflow_id = task
1741 .workflow_id
1742 .ok_or("task missing workflow id")
1743 .and_then(|id| WorkflowId::try_from(id).map_err(|_| "bad workflow id"))?;
1744 let activity_id = task
1745 .activity_id
1746 .map(ActivityId::from)
1747 .ok_or("task missing activity id")?;
1748 sink.complete_activity(ActivityCompletion {
1749 workflow_id,
1750 activity_id,
1751 run_id: None,
1752 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1753 ContentType::Json,
1754 br#"{"greeting":"hello"}"#.to_vec(),
1755 )),
1756 })
1757 .map_err(|error| error.to_string())
1758 });
1759
1760 let dispatcher = Arc::new(
1761 WorkerActivityDispatcher::new(registry.clone(), "default", test_tracker())
1762 .with_pending(pending),
1763 );
1764
1765 let result = tokio::spawn(futures::future::lazy(move |_| {
1767 dispatcher.dispatch(greet_request())
1768 }))
1769 .await
1770 .map_err(|error| error.to_string())?;
1771
1772 assert_eq!(result, Ok(r#"{"greeting":"hello"}"#.to_owned()));
1773 echo.await.map_err(|error| error.to_string())??;
1774 on_n1.deregister()?;
1775 Ok(())
1776 }
1777}