1use std::collections::BTreeMap;
4
5use aion_core::{ActivityError, ActivityId, Payload, RunId, WorkflowId};
6use aion_proto::{
7 ProtoActivityId, ProtoActivityResult, ProtoActivityTask, ProtoPayload, ProtoRunId,
8 ProtoWorkflowId, WireError, proto_activity_result,
9};
10
11use crate::error::ServerError;
12use crate::shutdown::DrainState;
13use crate::worker::envelope::{CompletionFences, CompletionToken, idempotency_key};
14use crate::worker::queue_service::declarations::QueueDeclarationSource;
15use crate::worker::queue_service::policy::QueueServiceConfig;
16use crate::worker::queue_service::state::QueueServiceState;
17use crate::worker::queue_service::taxonomy::{QueueServiceReason, ServiceAddress};
18use crate::worker::queue_service::wait::{
19 ServiceWait, clear_selection_miss, observe_selection_miss,
20};
21use crate::worker::registry::{ConnectedWorkerRegistry, WorkerMessage};
22use tracing::{Instrument, info_span};
23
24#[derive(Clone, Debug, Eq, PartialEq)]
26pub struct ScheduledActivity {
27 pub namespace: String,
30 pub task_queue: String,
34 pub activity_type: String,
37 pub node: Option<String>,
44 pub workflow_id: WorkflowId,
46 pub activity_id: ActivityId,
48 pub run_id: Option<RunId>,
50 pub input: Payload,
52 pub attempt: u32,
55 pub labels: BTreeMap<String, String>,
58}
59
60impl ScheduledActivity {
61 fn require_run_id(&self) -> Result<&RunId, ServerError> {
65 self.run_id.as_ref().ok_or_else(|| {
66 ServerError::worker_dispatch(
67 self.namespace.clone(),
68 self.activity_type.clone(),
69 "activity run id is missing; refusing unfenced external effect",
70 )
71 })
72 }
73
74 pub fn to_task(
81 &self,
82 completion_token: &CompletionToken,
83 ) -> Result<ProtoActivityTask, ServerError> {
84 let run_id = self.require_run_id()?;
85 Ok(ProtoActivityTask {
86 workflow_id: Some(ProtoWorkflowId::from(self.workflow_id.clone())),
87 activity_id: Some(ProtoActivityId::from(self.activity_id.clone())),
88 activity_type: self.activity_type.clone(),
89 input: Some(ProtoPayload::from(self.input.clone())),
90 attempt: self.attempt,
91 labels: self.labels.clone().into_iter().collect(),
92 run_id: Some(ProtoRunId::from(run_id.clone())),
93 completion_token: completion_token.as_str().to_owned(),
94 idempotency_key: idempotency_key(&self.workflow_id, run_id, &self.activity_id),
95 })
96 }
97}
98
99#[derive(Clone, Debug)]
101pub struct ActivityDispatcher {
102 registry: ConnectedWorkerRegistry,
103 drain_state: DrainState,
104 completion_fences: CompletionFences,
105 queue_declarations: QueueDeclarationSource,
114 queue_service_state: QueueServiceState,
115 queue_service_config: QueueServiceConfig,
116}
117
118impl ActivityDispatcher {
119 #[must_use]
121 pub fn new(registry: ConnectedWorkerRegistry) -> Self {
122 Self {
123 registry,
124 drain_state: DrainState::default(),
125 completion_fences: CompletionFences::default(),
126 queue_declarations: QueueDeclarationSource::default(),
127 queue_service_state: QueueServiceState::default(),
128 queue_service_config: QueueServiceConfig::default(),
129 }
130 }
131
132 #[must_use]
135 pub fn with_queue_service(
136 mut self,
137 declarations: QueueDeclarationSource,
138 state: QueueServiceState,
139 config: QueueServiceConfig,
140 ) -> Self {
141 self.queue_declarations = declarations;
142 self.queue_service_state = state;
143 self.queue_service_config = config;
144 self
145 }
146
147 #[must_use]
149 pub fn with_drain_state(mut self, drain_state: DrainState) -> Self {
150 self.drain_state = drain_state;
151 self
152 }
153
154 #[must_use]
156 pub fn with_completion_fences(mut self, completion_fences: CompletionFences) -> Self {
157 self.completion_fences = completion_fences;
158 self
159 }
160
161 pub async fn dispatch(&self, activity: &ScheduledActivity) -> Result<(), ServerError> {
168 let span = info_span!(
169 "activity_dispatch",
170 operation = "activity_dispatch",
171 namespace = %activity.namespace,
172 task_queue = %activity.task_queue,
173 node = activity.node.as_deref(),
174 workflow_id = %activity.workflow_id,
175 activity_id = %activity.activity_id,
176 activity_type = %activity.activity_type,
177 worker_id = tracing::field::Empty,
178 );
179 let span_fields = span.clone();
180
181 async {
182 self.dispatch_to_node(activity, activity.node.as_deref(), &span_fields)
183 .await
184 }
185 .instrument(span)
186 .await
187 .inspect_err(|error| {
188 log_dispatch_error("activity_dispatch", activity, error);
189 })
190 }
191
192 pub async fn dispatch_preferring(
220 &self,
221 activity: &ScheduledActivity,
222 preferred: &std::collections::BTreeSet<String>,
223 ) -> Result<(), ServerError> {
224 let tiers = crate::worker::preferred_node_order(&aion_store::NamespacePlacement::Prefer {
227 nodes: preferred.clone(),
228 });
229 self.dispatch_over_tiers(activity, &tiers).await
230 }
231
232 pub async fn dispatch_requiring(
260 &self,
261 activity: &ScheduledActivity,
262 required: &std::collections::BTreeSet<String>,
263 ) -> Result<(), ServerError> {
264 let span = info_span!(
265 "activity_dispatch",
266 operation = "activity_dispatch_requiring",
267 namespace = %activity.namespace,
268 task_queue = %activity.task_queue,
269 workflow_id = %activity.workflow_id,
270 activity_id = %activity.activity_id,
271 activity_type = %activity.activity_type,
272 worker_id = tracing::field::Empty,
273 );
274 let span_fields = span.clone();
275 async {
276 loop {
277 for label in required {
278 self.drain_state
279 .ensure_accepting(&activity.namespace, &activity.activity_type)?;
280 let candidates = self.registry.workers_for(
281 &activity.namespace,
282 &activity.task_queue,
283 &activity.activity_type,
284 Some(label.as_str()),
285 )?;
286 if let Some(()) = self
287 .send_to_candidates(activity, candidates, &span_fields)
288 .await?
289 {
290 return Ok(());
291 }
292 }
293 tracing::info!(
297 namespace = %activity.namespace,
298 task_queue = %activity.task_queue,
299 activity_type = %activity.activity_type,
300 workflow_id = %activity.workflow_id,
301 activity_id = %activity.activity_id,
302 "no worker on a required (Pinned) node; waiting — will NOT spill to any-node"
303 );
304 self.registry.wait_for_worker().await;
305 }
306 }
307 .instrument(span)
308 .await
309 .inspect_err(|error| {
310 log_dispatch_error("activity_dispatch_requiring", activity, error);
311 })
312 }
313
314 async fn dispatch_over_tiers(
326 &self,
327 activity: &ScheduledActivity,
328 tiers: &[Option<String>],
329 ) -> Result<(), ServerError> {
330 let span = info_span!(
331 "activity_dispatch",
332 operation = "activity_dispatch_preferring",
333 namespace = %activity.namespace,
334 task_queue = %activity.task_queue,
335 workflow_id = %activity.workflow_id,
336 activity_id = %activity.activity_id,
337 activity_type = %activity.activity_type,
338 worker_id = tracing::field::Empty,
339 );
340 let span_fields = span.clone();
341 async {
342 for tier in tiers {
343 let Some(label) = tier else {
344 return self
347 .dispatch_to_node(activity, activity.node.as_deref(), &span_fields)
348 .await;
349 };
350 self.drain_state
351 .ensure_accepting(&activity.namespace, &activity.activity_type)?;
352 let candidates = self.registry.workers_for(
353 &activity.namespace,
354 &activity.task_queue,
355 &activity.activity_type,
356 Some(label.as_str()),
357 )?;
358 if let Some(()) = self
359 .send_to_candidates(activity, candidates, &span_fields)
360 .await?
361 {
362 return Ok(());
363 }
364 }
365 self.dispatch_to_node(activity, activity.node.as_deref(), &span_fields)
368 .await
369 }
370 .instrument(span)
371 .await
372 .inspect_err(|error| {
373 log_dispatch_error("activity_dispatch_preferring", activity, error);
374 })
375 }
376
377 async fn dispatch_to_node(
380 &self,
381 activity: &ScheduledActivity,
382 node: Option<&str>,
383 span_fields: &tracing::Span,
384 ) -> Result<(), ServerError> {
385 let address = ServiceAddress {
394 namespace: activity.namespace.clone(),
395 task_queue: activity.task_queue.clone(),
396 activity_type: activity.activity_type.clone(),
397 node: node.map(ToOwned::to_owned),
398 };
399 let wait = ServiceWait {
400 registry: &self.registry,
401 declarations: &self.queue_declarations,
402 config: &self.queue_service_config,
403 state: &self.queue_service_state,
404 address: &address,
405 workflow_id: &activity.workflow_id,
406 activity_id: &activity.activity_id,
407 };
408 let policy = self
409 .queue_service_config
410 .policy_for(&activity.namespace, &activity.task_queue);
411 let started_at = std::time::Instant::now();
412 let mut reported: Option<QueueServiceReason> = None;
413 let workers = loop {
414 self.drain_state
415 .ensure_accepting(&activity.namespace, &activity.activity_type)
416 .inspect_err(|_| clear_selection_miss(&wait))?;
417 let candidates = self
418 .registry
419 .workers_for(
420 &activity.namespace,
421 &activity.task_queue,
422 &activity.activity_type,
423 node,
424 )
425 .inspect_err(|_| clear_selection_miss(&wait))?;
426 if !candidates.is_empty() {
427 if let Some(reason) = reported {
428 tracing::info!(
429 namespace = %activity.namespace,
430 task_queue = %activity.task_queue,
431 activity_type = %activity.activity_type,
432 workflow_id = %activity.workflow_id,
433 activity_id = %activity.activity_id,
434 queue_service_reason = reason.as_str(),
435 "queue service restored; the parked dispatch has a worker"
436 );
437 }
438 clear_selection_miss(&wait);
439 break candidates;
440 }
441 match observe_selection_miss(&wait, policy, None, started_at.elapsed(), reported) {
442 Ok(None) => continue,
445 Ok(Some(observed)) => reported = Some(observed.reason),
446 Err(refusal) => {
447 clear_selection_miss(&wait);
448 return Err(ServerError::worker_dispatch(
449 activity.namespace.clone(),
450 activity.activity_type.clone(),
451 refusal.reason_string(),
452 ));
453 }
454 }
455 self.registry.wait_for_worker().await;
456 };
457 match self
458 .send_to_candidates(activity, workers, span_fields)
459 .await?
460 {
461 Some(()) => Ok(()),
462 None => Err(ServerError::worker_dispatch(
463 activity.namespace.clone(),
464 activity.activity_type.clone(),
465 format!(
466 "all matching worker streams in task queue {} closed before task could be \
467 delivered",
468 activity.task_queue
469 ),
470 )),
471 }
472 }
473
474 async fn send_to_candidates(
479 &self,
480 activity: &ScheduledActivity,
481 candidates: Vec<crate::worker::registry::WorkerHandle>,
482 span_fields: &tracing::Span,
483 ) -> Result<Option<()>, ServerError> {
484 activity.require_run_id()?;
485 let completion_token = self
486 .completion_fences
487 .issue(&activity.workflow_id, &activity.activity_id)?;
488 let task = activity.to_task(&completion_token)?;
489 for worker in candidates {
490 if let Err(error) = self
491 .drain_state
492 .ensure_accepting(&activity.namespace, &activity.activity_type)
493 {
494 self.completion_fences.revoke(
495 &activity.workflow_id,
496 &activity.activity_id,
497 &completion_token,
498 )?;
499 return Err(error);
500 }
501 span_fields.record("worker_id", format!("{:?}", worker.id()));
502 if let Some(sender) = worker.sender() {
507 if sender
508 .send(WorkerMessage::ActivityTask(Box::new(task.clone())))
509 .await
510 .is_ok()
511 {
512 return Ok(Some(()));
513 }
514 }
515 self.registry.deregister(worker.id())?;
516 }
517 self.completion_fences.revoke(
518 &activity.workflow_id,
519 &activity.activity_id,
520 &completion_token,
521 )?;
522 Ok(None)
523 }
524}
525
526fn log_dispatch_error(operation: &'static str, activity: &ScheduledActivity, error: &ServerError) {
527 let fields = error.trace_fields();
528 tracing::error!(
529 operation,
530 namespace = %activity.namespace,
531 task_queue = %activity.task_queue,
532 node = activity.node.as_deref(),
533 workflow_id = %activity.workflow_id,
534 activity_id = %activity.activity_id,
535 activity_type = %activity.activity_type,
536 error_type = %fields.error_type,
537 store_error_type = fields.store_error_type,
538 reason = %fields.reason,
539 "activity dispatch failed"
540 );
541}
542
543#[derive(Clone, Debug, Eq, PartialEq)]
545pub enum ActivityCompletionOutcome {
546 Succeeded(Payload),
548 Failed(ActivityError),
550 WorkerLost {
561 worker_id: crate::worker::registry::WorkerId,
563 },
564}
565
566#[derive(Clone, Debug, Eq, PartialEq)]
568pub struct ActivityCompletion {
569 pub workflow_id: WorkflowId,
571 pub activity_id: ActivityId,
573 pub run_id: Option<RunId>,
575 pub completion_token: CompletionToken,
577 pub outcome: ActivityCompletionOutcome,
579}
580
581impl TryFrom<ProtoActivityResult> for ActivityCompletion {
582 type Error = ServerError;
583
584 fn try_from(value: ProtoActivityResult) -> Result<Self, Self::Error> {
585 let workflow_id = value
586 .workflow_id
587 .ok_or_else(|| wire_error("activity result workflow id is missing"))
588 .and_then(|id| WorkflowId::try_from(id).map_err(ServerError::from))?;
589 let activity_id = value
590 .activity_id
591 .ok_or_else(|| wire_error("activity result activity id is missing"))
592 .map(ActivityId::from)?;
593 let run_id = value
594 .run_id
595 .ok_or_else(|| wire_error("activity result run id is missing"))
596 .and_then(|id| RunId::try_from(id).map_err(ServerError::from))?;
597 let completion_token =
598 CompletionToken::from_wire(&workflow_id, &activity_id, value.completion_token)?;
599 let outcome = match value.outcome {
600 Some(proto_activity_result::Outcome::Result(payload)) => {
601 ActivityCompletionOutcome::Succeeded(
602 Payload::try_from(payload).map_err(ServerError::from)?,
603 )
604 }
605 Some(proto_activity_result::Outcome::Error(error)) => {
606 ActivityCompletionOutcome::Failed(
607 ActivityError::try_from(error).map_err(ServerError::from)?,
608 )
609 }
610 None => return Err(wire_error("activity result outcome is missing")),
611 };
612
613 Ok(Self {
614 workflow_id,
615 activity_id,
616 run_id: Some(run_id),
617 completion_token,
618 outcome,
619 })
620 }
621}
622
623pub trait ActivityCompletionSink {
625 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError>;
631
632 fn park_activity(
649 &self,
650 workflow_id: &WorkflowId,
651 activity_id: &ActivityId,
652 ) -> Result<(), ServerError>;
653}
654
655pub fn handle_activity_result(
661 sink: &impl ActivityCompletionSink,
662 result: ProtoActivityResult,
663) -> Result<(), ServerError> {
664 sink.complete_activity(ActivityCompletion::try_from(result)?)
665}
666
667fn wire_error(message: &'static str) -> ServerError {
668 ServerError::Wire {
669 wire: WireError::backend(message),
670 }
671}
672
673#[cfg(test)]
674mod tests {
675 use std::sync::Mutex;
676
677 use aion_core::{ActivityErrorKind, ContentType};
678 use aion_proto::{ProtoActivityError, ProtoActivityErrorKind};
679 use serde_json::json;
680 use uuid::Uuid;
681
682 use crate::worker::registry::ConnectedWorkerRegistry;
683
684 use super::*;
685
686 fn workflow_id() -> WorkflowId {
687 WorkflowId::new(Uuid::nil())
688 }
689
690 fn activity_id() -> ActivityId {
691 ActivityId::from_sequence_position(42)
692 }
693
694 fn payload(value: &serde_json::Value) -> Result<Payload, Box<dyn std::error::Error>> {
695 Ok(Payload::from_json(value)?)
696 }
697
698 #[tokio::test]
699 async fn dispatch_pushes_activity_task_with_correlation()
700 -> Result<(), Box<dyn std::error::Error>> {
701 let registry = ConnectedWorkerRegistry::default();
702 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
703 let activity_types = [String::from("charge-card")];
704 let registration = registry.register("tenant-a", activity_types.iter(), tx)?;
705 let dispatcher = ActivityDispatcher::new(registry.clone());
706 let input = payload(&json!({"amount": 1200}))?;
707 let scheduled = ScheduledActivity {
708 namespace: String::from("tenant-a"),
709 task_queue: String::from("default"),
710 activity_type: String::from("charge-card"),
711 node: None,
712 workflow_id: workflow_id(),
713 activity_id: activity_id(),
714 run_id: Some(RunId::new_v4()),
715 input: input.clone(),
716 attempt: 1,
717 labels: std::collections::BTreeMap::new(),
718 };
719
720 dispatcher.dispatch(&scheduled).await?;
721 let message = rx.recv().await.ok_or("expected pushed activity task")?;
722 let WorkerMessage::ActivityTask(task) = message else {
723 return Err("expected activity task message".into());
724 };
725
726 assert_eq!(task.workflow_id, Some(ProtoWorkflowId::from(workflow_id())));
727 assert_eq!(task.activity_id, Some(ProtoActivityId::from(activity_id())));
728 assert_eq!(task.activity_type, "charge-card");
729 assert_eq!(task.input, Some(ProtoPayload::from(input)));
730 assert_eq!(task.attempt, 1, "wire task must carry the stamped attempt");
731
732 registration.deregister()?;
733 Ok(())
734 }
735
736 #[tokio::test]
737 async fn dispatch_waits_for_worker_then_delivers() -> Result<(), Box<dyn std::error::Error>> {
738 let registry = ConnectedWorkerRegistry::default();
739 let dispatcher = ActivityDispatcher::new(registry.clone());
740 let scheduled = ScheduledActivity {
741 namespace: String::from("tenant-a"),
742 task_queue: String::from("default"),
743 activity_type: String::from("charge-card"),
744 node: None,
745 workflow_id: workflow_id(),
746 activity_id: activity_id(),
747 run_id: Some(RunId::new_v4()),
748 input: Payload::new(ContentType::Json, b"{}".to_vec()),
749 attempt: 1,
750 labels: std::collections::BTreeMap::new(),
751 };
752
753 let dispatch_handle = tokio::spawn({
754 let dispatcher = dispatcher.clone();
755 let scheduled = scheduled.clone();
756 async move { dispatcher.dispatch(&scheduled).await }
757 });
758
759 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
760 assert!(!dispatch_handle.is_finished(), "dispatch should be waiting");
761
762 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
763 let activity_types = [String::from("charge-card")];
764 let _registration = registry.register("tenant-a", activity_types.iter(), tx)?;
765
766 dispatch_handle.await??;
767 assert!(rx.recv().await.is_some());
768 Ok(())
769 }
770
771 #[tokio::test]
783 async fn a_dispatch_with_no_worker_publishes_its_park_and_clears_on_arrival()
784 -> Result<(), Box<dyn std::error::Error>> {
785 let registry = ConnectedWorkerRegistry::default();
786 let state = QueueServiceState::default();
787 let dispatcher = ActivityDispatcher::new(registry.clone()).with_queue_service(
788 QueueDeclarationSource::default(),
789 state.clone(),
790 QueueServiceConfig::default(),
791 );
792 let scheduled = ScheduledActivity {
793 namespace: String::from("tenant-a"),
794 task_queue: String::from("default"),
795 activity_type: String::from("charge-card"),
796 node: None,
797 workflow_id: workflow_id(),
798 activity_id: activity_id(),
799 run_id: Some(RunId::new_v4()),
800 input: Payload::new(ContentType::Json, b"{}".to_vec()),
801 attempt: 1,
802 labels: std::collections::BTreeMap::new(),
803 };
804
805 assert!(
808 state.unserved()?.is_empty(),
809 "no dispatch has been made yet"
810 );
811
812 let unwired = ActivityDispatcher::new(registry.clone());
818 let unwired_handle = tokio::spawn({
819 let scheduled = scheduled.clone();
820 async move { unwired.dispatch(&scheduled).await }
821 });
822 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
823 assert!(
824 state.unserved()?.is_empty(),
825 "an unshared dispatcher must publish nothing HERE — that is the \
826 defect this test exists to catch, reproduced on purpose"
827 );
828 unwired_handle.abort();
829
830 let dispatch_handle = tokio::spawn({
831 let dispatcher = dispatcher.clone();
832 let scheduled = scheduled.clone();
833 async move { dispatcher.dispatch(&scheduled).await }
834 });
835 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
836 assert!(!dispatch_handle.is_finished(), "dispatch should be waiting");
837
838 let unserved = state.unserved()?;
839 assert_eq!(
840 unserved.len(),
841 1,
842 "the parked dispatch must be queryable, not just logged: {unserved:?}"
843 );
844 assert_eq!(unserved[0].key.task_queue, "default");
845 assert_eq!(
846 unserved[0].reason,
847 QueueServiceReason::NoLivePollers,
848 "an empty pool must be classified, not reported as a bare miss"
849 );
850 assert_eq!(
851 state.parked_on_queue("default")?,
852 1,
853 "the run parked on the queue must be attributable to the queue"
854 );
855
856 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
859 let activity_types = [String::from("charge-card")];
860 let _registration = registry.register("tenant-a", activity_types.iter(), tx)?;
861
862 dispatch_handle.await??;
863 assert!(rx.recv().await.is_some(), "the task must be delivered");
864 assert!(
865 state.unserved()?.is_empty(),
866 "a served dispatch must not be left published as unserved: {:?}",
867 state.unserved()?
868 );
869 Ok(())
870 }
871
872 #[tokio::test]
873 async fn dispatch_skips_closed_worker_and_uses_next_match()
874 -> Result<(), Box<dyn std::error::Error>> {
875 let registry = ConnectedWorkerRegistry::default();
876 let (closed_tx, closed_rx) = tokio::sync::mpsc::channel(1);
877 let (live_tx, mut live_rx) = tokio::sync::mpsc::channel(1);
878 let activity_types = [String::from("charge-card")];
879 let closed_registration =
880 registry.register("tenant-a", activity_types.iter(), closed_tx)?;
881 let live_registration = registry.register("tenant-a", activity_types.iter(), live_tx)?;
882 drop(closed_rx);
883
884 let dispatcher = ActivityDispatcher::new(registry.clone());
885 let scheduled = ScheduledActivity {
886 namespace: String::from("tenant-a"),
887 task_queue: String::from("default"),
888 activity_type: String::from("charge-card"),
889 node: None,
890 workflow_id: workflow_id(),
891 activity_id: activity_id(),
892 run_id: Some(RunId::new_v4()),
893 input: Payload::new(ContentType::Json, b"{}".to_vec()),
894 attempt: 1,
895 labels: std::collections::BTreeMap::new(),
896 };
897
898 dispatcher.dispatch(&scheduled).await?;
899
900 assert!(live_rx.recv().await.is_some());
901 assert_eq!(
902 registry
903 .workers_for("tenant-a", "default", "charge-card", None)?
904 .len(),
905 1
906 );
907
908 closed_registration.deregister()?;
909 live_registration.deregister()?;
910 Ok(())
911 }
912
913 fn scheduled_unpinned() -> ScheduledActivity {
914 ScheduledActivity {
915 namespace: String::from("tenant-a"),
916 task_queue: String::from("default"),
917 activity_type: String::from("charge-card"),
918 node: None,
921 workflow_id: workflow_id(),
922 activity_id: activity_id(),
923 run_id: Some(RunId::new_v4()),
924 input: Payload::new(ContentType::Json, b"{}".to_vec()),
925 attempt: 1,
926 labels: std::collections::BTreeMap::new(),
927 }
928 }
929
930 fn required(labels: &[&str]) -> std::collections::BTreeSet<String> {
931 labels.iter().map(|l| (*l).to_owned()).collect()
932 }
933
934 #[tokio::test]
939 async fn dispatch_requiring_waits_and_never_spills_to_a_wrong_node_worker()
940 -> Result<(), Box<dyn std::error::Error>> {
941 let registry = ConnectedWorkerRegistry::default();
942 let dispatcher = ActivityDispatcher::new(registry.clone());
943 let scheduled = scheduled_unpinned();
944 let types = [String::from("charge-card")];
945
946 let (wrong_tx, mut wrong_rx) = tokio::sync::mpsc::channel(1);
949 let _wrong = registry.register_namespaces(
950 [String::from("tenant-a")],
951 "default",
952 Some(String::from("n2")),
953 types.iter(),
954 wrong_tx,
955 )?;
956
957 let handle = tokio::spawn({
958 let dispatcher = dispatcher.clone();
959 let scheduled = scheduled.clone();
960 async move {
961 dispatcher
962 .dispatch_requiring(&scheduled, &required(&["n1"]))
963 .await
964 }
965 });
966
967 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
969 assert!(
970 !handle.is_finished(),
971 "Pinned{{n1}} must WAIT rather than spill to the live n2 worker"
972 );
973 assert!(
974 wrong_rx.try_recv().is_err(),
975 "the wrong-node (n2) worker must never receive the task"
976 );
977
978 let (right_tx, mut right_rx) = tokio::sync::mpsc::channel(1);
980 let _right = registry.register_namespaces(
981 [String::from("tenant-a")],
982 "default",
983 Some(String::from("n1")),
984 types.iter(),
985 right_tx,
986 )?;
987
988 handle.await??;
989 assert!(
990 right_rx.recv().await.is_some(),
991 "the required n1 worker receives the task once live"
992 );
993 assert!(
994 wrong_rx.try_recv().is_err(),
995 "the wrong-node worker still never received it"
996 );
997 Ok(())
998 }
999
1000 #[tokio::test]
1003 async fn dispatch_requiring_never_mutates_the_rows_node()
1004 -> Result<(), Box<dyn std::error::Error>> {
1005 let registry = ConnectedWorkerRegistry::default();
1006 let dispatcher = ActivityDispatcher::new(registry.clone());
1007 let scheduled = scheduled_unpinned();
1008 assert_eq!(scheduled.node, None, "precondition: the row is unpinned");
1009 let types = [String::from("charge-card")];
1010 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
1011 let _right = registry.register_namespaces(
1012 [String::from("tenant-a")],
1013 "default",
1014 Some(String::from("n1")),
1015 types.iter(),
1016 tx,
1017 )?;
1018
1019 dispatcher
1020 .dispatch_requiring(&scheduled, &required(&["n1"]))
1021 .await?;
1022
1023 assert!(rx.recv().await.is_some(), "the n1 worker received the task");
1024 assert_eq!(
1025 scheduled.node, None,
1026 "the row's authored node MUST remain None through a Pinned dispatch \
1027 (the determinism invariant, CP-Phase-2 §2.4)"
1028 );
1029 Ok(())
1030 }
1031
1032 #[derive(Default)]
1033 struct RecordingSink {
1034 completions: Mutex<Vec<ActivityCompletion>>,
1035 }
1036
1037 impl ActivityCompletionSink for RecordingSink {
1038 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
1039 self.completions
1040 .lock()
1041 .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
1042 .push(completion);
1043 Ok(())
1044 }
1045
1046 fn park_activity(
1047 &self,
1048 _workflow_id: &WorkflowId,
1049 _activity_id: &ActivityId,
1050 ) -> Result<(), ServerError> {
1051 Err(ServerError::worker_dispatch(
1052 "",
1053 "",
1054 "result-handoff tests never park a dispatch",
1055 ))
1056 }
1057 }
1058
1059 #[test]
1060 fn successful_activity_result_calls_completion_sink() -> Result<(), Box<dyn std::error::Error>>
1061 {
1062 let sink = RecordingSink::default();
1063 let output = payload(&json!({"ok": true}))?;
1064 let result = ProtoActivityResult {
1065 workflow_id: Some(ProtoWorkflowId::from(workflow_id())),
1066 activity_id: Some(ProtoActivityId::from(activity_id())),
1067 run_id: Some(ProtoRunId::from(RunId::new_v4())),
1068 completion_token: String::from("generation-1"),
1069 outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
1070 output.clone(),
1071 ))),
1072 };
1073
1074 handle_activity_result(&sink, result)?;
1075 let completions = sink
1076 .completions
1077 .lock()
1078 .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
1079
1080 assert_eq!(completions.len(), 1);
1081 assert_eq!(completions[0].workflow_id, workflow_id());
1082 assert_eq!(completions[0].activity_id, activity_id());
1083 assert_eq!(
1084 completions[0].outcome,
1085 ActivityCompletionOutcome::Succeeded(output)
1086 );
1087 Ok(())
1088 }
1089
1090 #[test]
1091 fn failed_activity_result_preserves_error_classification()
1092 -> Result<(), Box<dyn std::error::Error>> {
1093 let sink = RecordingSink::default();
1094 let error = ProtoActivityError {
1095 kind: ProtoActivityErrorKind::Retryable as i32,
1096 message: String::from("temporary outage"),
1097 details: Some(ProtoPayload::from(payload(
1098 &json!({"retry_after_ms": 500}),
1099 )?)),
1100 };
1101 let result = ProtoActivityResult {
1102 workflow_id: Some(ProtoWorkflowId::from(workflow_id())),
1103 activity_id: Some(ProtoActivityId::from(activity_id())),
1104 run_id: Some(ProtoRunId::from(RunId::new_v4())),
1105 completion_token: String::from("generation-1"),
1106 outcome: Some(proto_activity_result::Outcome::Error(error)),
1107 };
1108
1109 handle_activity_result(&sink, result)?;
1110 let completions = sink
1111 .completions
1112 .lock()
1113 .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
1114
1115 assert_eq!(completions.len(), 1);
1116 match &completions[0].outcome {
1117 ActivityCompletionOutcome::Failed(error) => {
1118 assert_eq!(error.kind, ActivityErrorKind::Retryable);
1119 assert!(error.is_retryable());
1120 }
1121 other => return Err(format!("expected failed outcome, got {other:?}").into()),
1122 }
1123 Ok(())
1124 }
1125}