1use anyhow::bail;
7use parking_lot::RwLock;
8use rand::seq::SliceRandom;
9use std::{
10 collections::{
11 HashMap,
12 hash_map::Entry::{Occupied, Vacant},
13 },
14 future::Future,
15 sync::{Arc, Weak},
16};
17use temporalio_common::{
18 protos::{
19 TaskToken,
20 temporal::api::{
21 worker::v1::WorkerHeartbeat,
22 workflowservice::v1::{DescribeNamespaceResponse, PollWorkflowTaskQueueResponse},
23 },
24 },
25 worker::{WorkerDeploymentOptions, WorkerTaskTypes},
26};
27use tokio::sync::OnceCell;
28use tonic::Code;
29use uuid::Uuid;
30
31#[cfg_attr(test, mockall::automock)]
33pub trait Slot {
34 fn schedule_wft(
36 self: Box<Self>,
37 task: PollWorkflowTaskQueueResponse,
38 ) -> Result<(), anyhow::Error>;
39}
40
41pub(crate) struct SlotReservation {
43 pub slot: Box<dyn Slot + Send>,
45 pub deployment_options: Option<WorkerDeploymentOptions>,
47}
48
49#[derive(PartialEq, Eq, Hash, Debug, Clone)]
50struct SlotKey {
51 namespace: String,
52 task_queue: String,
53}
54
55impl SlotKey {
56 fn new(namespace: String, task_queue: String) -> SlotKey {
57 SlotKey {
58 namespace,
59 task_queue,
60 }
61 }
62}
63
64#[derive(Debug, Clone)]
66struct RegisteredWorkerInfo {
67 worker_id: Uuid,
69 build_id: Option<String>,
71 task_types: WorkerTaskTypes,
73}
74
75impl RegisteredWorkerInfo {
76 fn new(worker_id: Uuid, build_id: Option<String>, task_types: WorkerTaskTypes) -> Self {
77 Self {
78 worker_id,
79 build_id,
80 task_types,
81 }
82 }
83}
84
85struct ClientWorkerSetImpl {
87 slot_providers: HashMap<SlotKey, Vec<RegisteredWorkerInfo>>,
89 all_workers: HashMap<Uuid, Arc<dyn ClientWorker + Send + Sync>>,
91 shared_worker: HashMap<String, Box<dyn SharedNamespaceWorkerTrait + Send + Sync>>,
93 namespace_descriptions: HashMap<String, Weak<NamespaceDescriptionSource>>,
95}
96
97impl ClientWorkerSetImpl {
98 fn new() -> Self {
100 Self {
101 slot_providers: Default::default(),
102 all_workers: Default::default(),
103 shared_worker: Default::default(),
104 namespace_descriptions: Default::default(),
105 }
106 }
107
108 fn namespace_description_source(&mut self, namespace: &str) -> Arc<NamespaceDescriptionSource> {
109 if let Some(description) = self
110 .namespace_descriptions
111 .get(namespace)
112 .and_then(Weak::upgrade)
113 {
114 return description;
115 }
116
117 let description = Arc::new(NamespaceDescriptionSource::unresolved());
118 self.namespace_descriptions
119 .insert(namespace.to_owned(), Arc::downgrade(&description));
120 description
121 }
122
123 fn try_reserve_wft_slot(
124 &self,
125 namespace: String,
126 task_queue: String,
127 ) -> Option<SlotReservation> {
128 let key = SlotKey::new(namespace, task_queue);
129 if let Some(worker_list) = self.slot_providers.get(&key) {
130 let workflow_workers: Vec<&RegisteredWorkerInfo> = worker_list
131 .iter()
132 .filter(|info| info.task_types.enable_workflows)
133 .collect();
134
135 for worker_id in Self::worker_ids_in_selection_order(&workflow_workers) {
136 if let Some(worker) = self.all_workers.get(&worker_id)
137 && let Some(slot) = worker.try_reserve_wft_slot()
138 {
139 let deployment_options = worker.deployment_options();
140 return Some(SlotReservation {
141 slot,
142 deployment_options,
143 });
144 }
145 }
146 }
147 None
148 }
149
150 fn worker_control_task_queue_enabled(&self, namespace: &str) -> bool {
151 self.shared_worker
152 .get(namespace)
153 .is_some_and(|worker| worker.worker_control_task_queue_enabled())
154 }
155
156 fn worker_ids_in_selection_order(worker_list: &[&RegisteredWorkerInfo]) -> Vec<Uuid> {
157 if cfg!(test) {
160 worker_list.iter().map(|info| info.worker_id).collect()
161 } else {
162 let mut rng = rand::rng();
163 let mut shuffled: Vec<_> = worker_list.to_vec();
164 shuffled.shuffle(&mut rng);
165 shuffled.iter().map(|info| info.worker_id).collect()
166 }
167 }
168
169 fn register(
170 &mut self,
171 worker: Arc<dyn ClientWorker + Send + Sync>,
172 skip_client_worker_set_check: bool,
173 ) -> Result<(), anyhow::Error> {
174 let slot_key = SlotKey::new(
175 worker.namespace().to_string(),
176 worker.task_queue().to_string(),
177 );
178 let build_id = worker
179 .deployment_options()
180 .map(|opts| opts.version.build_id);
181 let task_types = worker.worker_task_types();
182
183 if !task_types.enable_workflows
184 && !task_types.enable_local_activities
185 && !task_types.enable_remote_activities
186 && !task_types.enable_nexus
187 {
188 bail!(
189 "Worker must have at least one capability enabled (workflows, activities, or nexus)"
190 );
191 }
192
193 if !task_types.enable_workflows && task_types.enable_local_activities {
194 bail!("Local activities cannot be enabled without workflows")
195 }
196
197 if !skip_client_worker_set_check
198 && let Some(existing_workers) = self.slot_providers.get(&slot_key)
199 {
200 for existing_worker_info in existing_workers {
201 if existing_worker_info.build_id.as_ref() == build_id.as_ref()
202 && task_types.overlaps_with(&existing_worker_info.task_types)
203 {
204 bail!(
205 "Registration of multiple workers with overlapping worker task types \
206 on the same namespace, task queue, and deployment build ID not allowed: \
207 {slot_key:?}, worker_instance_key: {:?} \
208 build_id: {build_id:?}, \
209 new task types: {task_types:?}, \
210 existing task types: {:?}.",
211 existing_worker_info.task_types,
212 worker.worker_instance_key()
213 );
214 }
215 }
216 }
217
218 if worker.heartbeat_enabled()
219 && let Some(heartbeat_callback) = worker.heartbeat_callback()
220 {
221 let worker_instance_key = worker.worker_instance_key();
222 let namespace = worker.namespace().to_string();
223
224 let shared_worker = match self.shared_worker.entry(namespace) {
225 Occupied(o) => o.into_mut(),
226 Vacant(v) => {
227 let shared_worker = worker.new_shared_namespace_worker()?;
228 v.insert(shared_worker)
229 }
230 };
231 shared_worker.register_callback(
232 worker_instance_key,
233 WorkerCallbacks::new(
234 heartbeat_callback,
235 worker.heartbeat_success_callback(),
236 worker.cancel_activity_callback(),
237 ),
238 );
239 }
240
241 let worker_info =
242 RegisteredWorkerInfo::new(worker.worker_instance_key(), build_id, task_types);
243
244 match self.slot_providers.entry(slot_key.clone()) {
245 Occupied(o) => o.into_mut().push(worker_info),
246 Vacant(v) => {
247 v.insert(vec![worker_info]);
248 }
249 };
250
251 self.all_workers
252 .insert(worker.worker_instance_key(), worker);
253
254 Ok(())
255 }
256
257 fn unregister_slot_provider(&mut self, worker_instance_key: Uuid) -> Result<(), anyhow::Error> {
260 let worker = self.all_workers.get(&worker_instance_key).ok_or_else(|| {
261 anyhow::anyhow!("Worker not in all_workers during slot provider unregister")
262 })?;
263
264 let slot_key = SlotKey::new(
265 worker.namespace().to_string(),
266 worker.task_queue().to_string(),
267 );
268 if let Some(slot_vec) = self.slot_providers.get_mut(&slot_key) {
269 slot_vec.retain(|info| info.worker_id != worker_instance_key);
270 if slot_vec.is_empty() {
271 self.slot_providers.remove(&slot_key);
272 }
273 }
274 Ok(())
275 }
276
277 fn finalize_unregister(
278 &mut self,
279 worker_instance_key: Uuid,
280 ) -> Result<Arc<dyn ClientWorker + Send + Sync>, anyhow::Error> {
281 if let Some(worker) = self.all_workers.get(&worker_instance_key)
282 && let Some(slot_vec) = self.slot_providers.get(&SlotKey::new(
283 worker.namespace().to_string(),
284 worker.task_queue().to_string(),
285 ))
286 && slot_vec
287 .iter()
288 .any(|info| info.worker_id == worker_instance_key)
289 {
290 return Err(anyhow::anyhow!(
291 "Worker still in slot_providers during finalize"
292 ));
293 }
294
295 let worker = self
296 .all_workers
297 .remove(&worker_instance_key)
298 .ok_or_else(|| anyhow::anyhow!("Worker not found in all_workers"))?;
299
300 if let Some(w) = self.shared_worker.get_mut(worker.namespace()) {
301 let (callback, is_empty) = w.unregister_callback(worker.worker_instance_key());
302 if callback.is_some() && is_empty {
303 self.shared_worker.remove(worker.namespace());
304 }
305 }
306
307 Ok(worker)
308 }
309
310 #[cfg(test)]
311 fn num_providers(&self) -> usize {
312 self.slot_providers.values().map(|v| v.len()).sum()
313 }
314
315 #[cfg(test)]
316 fn num_heartbeat_workers(&self) -> usize {
317 self.shared_worker.values().map(|v| v.num_workers()).sum()
318 }
319}
320
321#[derive(Debug)]
323#[doc(hidden)]
324pub struct NamespaceDescriptionSource {
325 description: OnceCell<DescribeNamespaceResponse>,
326}
327
328impl NamespaceDescriptionSource {
329 pub fn unresolved() -> Self {
331 Self {
332 description: OnceCell::new(),
333 }
334 }
335
336 pub fn resolved(description: DescribeNamespaceResponse) -> Self {
338 Self {
339 description: OnceCell::new_with(Some(description)),
340 }
341 }
342
343 pub async fn resolve<F, Fut>(
345 &self,
346 fetch: F,
347 ) -> Result<&DescribeNamespaceResponse, tonic::Status>
348 where
349 F: FnOnce() -> Fut,
350 Fut: Future<Output = Result<DescribeNamespaceResponse, tonic::Status>>,
351 {
352 self.description
353 .get_or_try_init(|| async {
354 match fetch().await {
355 Err(status) if status.code() == Code::Unimplemented => {
356 Ok(DescribeNamespaceResponse::default())
357 }
358 result => result,
359 }
360 })
361 .await
362 }
363
364 pub fn get(&self) -> Option<&DescribeNamespaceResponse> {
366 self.description.get()
367 }
368}
369
370pub trait SharedNamespaceWorkerTrait {
373 fn namespace(&self) -> String;
375
376 fn register_callback(&self, worker_instance_key: Uuid, callbacks: WorkerCallbacks);
378
379 fn unregister_callback(&self, worker_instance_key: Uuid) -> (Option<WorkerCallbacks>, bool);
383
384 fn num_workers(&self) -> usize;
386
387 fn worker_control_task_queue_enabled(&self) -> bool {
389 false
390 }
391}
392
393pub struct ClientWorkerSet {
399 worker_grouping_key: Uuid,
400 worker_manager: RwLock<ClientWorkerSetImpl>,
401}
402
403impl Default for ClientWorkerSet {
404 fn default() -> Self {
405 Self::new()
406 }
407}
408
409impl ClientWorkerSet {
410 pub fn new() -> Self {
412 Self {
413 worker_grouping_key: Uuid::new_v4(),
414 worker_manager: RwLock::new(ClientWorkerSetImpl::new()),
415 }
416 }
417
418 #[doc(hidden)]
420 pub fn namespace_description_source(&self, namespace: &str) -> Arc<NamespaceDescriptionSource> {
421 self.worker_manager
422 .write()
423 .namespace_description_source(namespace)
424 }
425
426 pub(crate) fn try_reserve_wft_slot(
429 &self,
430 namespace: String,
431 task_queue: String,
432 ) -> Option<SlotReservation> {
433 self.worker_manager
434 .read()
435 .try_reserve_wft_slot(namespace, task_queue)
436 }
437
438 pub fn register_worker(
440 &self,
441 worker: Arc<dyn ClientWorker + Send + Sync>,
442 skip_client_worker_set_check: bool,
443 ) -> Result<(), anyhow::Error> {
444 self.worker_manager
445 .write()
446 .register(worker, skip_client_worker_set_check)
447 }
448
449 pub fn unregister_slot_provider(&self, worker_instance_key: Uuid) -> Result<(), anyhow::Error> {
452 self.worker_manager
453 .write()
454 .unregister_slot_provider(worker_instance_key)
455 }
456
457 pub fn finalize_unregister(
461 &self,
462 worker_instance_key: Uuid,
463 ) -> Result<Arc<dyn ClientWorker + Send + Sync>, anyhow::Error> {
464 self.worker_manager
465 .write()
466 .finalize_unregister(worker_instance_key)
467 }
468
469 pub fn worker_grouping_key(&self) -> Uuid {
471 self.worker_grouping_key
472 }
473
474 pub fn worker_control_task_queue_enabled(&self, namespace: &str) -> bool {
476 self.worker_manager
477 .read()
478 .worker_control_task_queue_enabled(namespace)
479 }
480
481 #[cfg(test)]
482 pub fn num_providers(&self) -> usize {
485 self.worker_manager.read().num_providers()
486 }
487
488 #[cfg(test)]
489 pub fn num_heartbeat_workers(&self) -> usize {
491 self.worker_manager.read().num_heartbeat_workers()
492 }
493}
494
495impl std::fmt::Debug for ClientWorkerSet {
496 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
497 f.debug_struct("ClientWorkerSet")
498 .field("worker_grouping_key", &self.worker_grouping_key)
499 .finish()
500 }
501}
502
503pub type HeartbeatCallback = Arc<dyn Fn() -> Option<WorkerHeartbeat> + Send + Sync>;
506
507pub type HeartbeatSuccessCallback = Arc<dyn Fn() + Send + Sync>;
509
510pub type CancelActivityCallback = Arc<dyn Fn(TaskToken) -> bool + Send + Sync>;
512
513#[non_exhaustive]
515pub struct WorkerCallbacks {
516 pub heartbeat: HeartbeatCallback,
518 pub heartbeat_success: Option<HeartbeatSuccessCallback>,
520 pub cancel_activity: Option<CancelActivityCallback>,
522}
523
524impl WorkerCallbacks {
525 pub fn new(
527 heartbeat: HeartbeatCallback,
528 heartbeat_success: Option<HeartbeatSuccessCallback>,
529 cancel_activity: Option<CancelActivityCallback>,
530 ) -> Self {
531 Self {
532 heartbeat,
533 heartbeat_success,
534 cancel_activity,
535 }
536 }
537}
538
539#[cfg_attr(test, mockall::automock)]
542pub trait ClientWorker: Send + Sync {
543 fn namespace(&self) -> &str;
545
546 fn task_queue(&self) -> &str;
548
549 fn try_reserve_wft_slot(&self) -> Option<Box<dyn Slot + Send>>;
555
556 fn deployment_options(&self) -> Option<WorkerDeploymentOptions>;
558
559 fn worker_instance_key(&self) -> Uuid;
562
563 fn heartbeat_enabled(&self) -> bool;
565
566 fn heartbeat_callback(&self) -> Option<HeartbeatCallback>;
568
569 fn heartbeat_success_callback(&self) -> Option<HeartbeatSuccessCallback> {
571 None
572 }
573
574 fn cancel_activity_callback(&self) -> Option<CancelActivityCallback>;
576
577 fn new_shared_namespace_worker(
579 &self,
580 ) -> Result<Box<dyn SharedNamespaceWorkerTrait + Send + Sync>, anyhow::Error>;
581
582 fn worker_task_types(&self) -> WorkerTaskTypes;
584}
585
586#[cfg(test)]
587mod tests {
588 use super::*;
589 use std::sync::atomic::{AtomicUsize, Ordering};
590
591 #[tokio::test]
592 async fn namespace_description_source_resolves_once() {
593 let source = NamespaceDescriptionSource::unresolved();
594 let calls = AtomicUsize::new(0);
595
596 let first = source.resolve(|| async {
597 calls.fetch_add(1, Ordering::Relaxed);
598 tokio::task::yield_now().await;
599 Ok(DescribeNamespaceResponse::default())
600 });
601 let second = source.resolve(|| async {
602 calls.fetch_add(1, Ordering::Relaxed);
603 Ok(DescribeNamespaceResponse::default())
604 });
605
606 let (first, second) = tokio::join!(first, second);
607 first.unwrap();
608 second.unwrap();
609 assert_eq!(calls.load(Ordering::Relaxed), 1);
610 }
611
612 #[tokio::test]
613 async fn namespace_description_source_resolves_unimplemented_as_default() {
614 let source = NamespaceDescriptionSource::unresolved();
615
616 let description = source
617 .resolve(|| async { Err(tonic::Status::unimplemented("unsupported")) })
618 .await
619 .unwrap();
620
621 assert_eq!(description, &DescribeNamespaceResponse::default());
622 }
623
624 #[test]
625 fn namespace_description_sources_are_scoped_by_namespace() {
626 let workers = ClientWorkerSet::new();
627 let first = workers.namespace_description_source("first");
628
629 assert!(Arc::ptr_eq(
630 &first,
631 &workers.namespace_description_source("first")
632 ));
633 assert!(!Arc::ptr_eq(
634 &first,
635 &workers.namespace_description_source("second")
636 ));
637 }
638
639 fn new_mock_slot(with_error: bool) -> Box<MockSlot> {
640 let mut mock_slot = MockSlot::new();
641 if with_error {
642 mock_slot
643 .expect_schedule_wft()
644 .returning(|_| Err(anyhow::anyhow!("Changed my mind")));
645 } else {
646 mock_slot.expect_schedule_wft().returning(|_| Ok(()));
647 }
648 Box::new(mock_slot)
649 }
650
651 fn new_mock_provider(
652 namespace: String,
653 task_queue: String,
654 with_error: bool,
655 no_slots: bool,
656 heartbeat_enabled: bool,
657 ) -> MockClientWorker {
658 let mut mock_provider = MockClientWorker::new();
659 mock_provider
660 .expect_try_reserve_wft_slot()
661 .returning(move || {
662 if no_slots {
663 None
664 } else {
665 Some(new_mock_slot(with_error))
666 }
667 });
668 mock_provider.expect_namespace().return_const(namespace);
669 mock_provider.expect_task_queue().return_const(task_queue);
670 mock_provider.expect_deployment_options().return_const(None);
671 mock_provider
672 .expect_heartbeat_enabled()
673 .return_const(heartbeat_enabled);
674 mock_provider
675 .expect_worker_instance_key()
676 .return_const(Uuid::new_v4());
677 mock_provider
678 .expect_worker_task_types()
679 .return_const(WorkerTaskTypes {
680 enable_workflows: true,
681 enable_local_activities: true,
682 enable_remote_activities: true,
683 enable_nexus: true,
684 });
685 mock_provider
686 }
687
688 #[test]
689 fn reserve_wft_slot_retries_another_worker_when_first_has_no_slot() {
690 let mut manager = ClientWorkerSetImpl::new();
691 let namespace = "retry_namespace".to_string();
692 let task_queue = "retry_queue".to_string();
693
694 let failing_worker_id = Uuid::new_v4();
695 let mut failing_worker = MockClientWorker::new();
696 failing_worker
697 .expect_try_reserve_wft_slot()
698 .times(1)
699 .returning(|| None);
700 failing_worker
701 .expect_namespace()
702 .return_const(namespace.clone());
703 failing_worker
704 .expect_task_queue()
705 .return_const(task_queue.clone());
706 failing_worker.expect_deployment_options().return_const(
707 WorkerDeploymentOptions::new(
708 temporalio_common::worker::WorkerDeploymentVersion::builder()
709 .deployment_name("test-deployment".to_string())
710 .build_id("build-fail".to_string())
711 .build(),
712 )
713 .use_worker_versioning(true)
714 .build(),
715 );
716 failing_worker
717 .expect_worker_instance_key()
718 .return_const(failing_worker_id);
719 failing_worker
720 .expect_heartbeat_enabled()
721 .return_const(false);
722 failing_worker
723 .expect_worker_task_types()
724 .return_const(WorkerTaskTypes {
725 enable_workflows: true,
726 enable_local_activities: true,
727 enable_remote_activities: true,
728 enable_nexus: true,
729 });
730
731 let succeeding_worker_id = Uuid::new_v4();
732 let mut succeeding_worker = MockClientWorker::new();
733 succeeding_worker
734 .expect_try_reserve_wft_slot()
735 .times(1)
736 .returning(|| Some(new_mock_slot(false)));
737 succeeding_worker
738 .expect_namespace()
739 .return_const(namespace.clone());
740 succeeding_worker
741 .expect_task_queue()
742 .return_const(task_queue.clone());
743 let success_deployment_options = WorkerDeploymentOptions::new(
744 temporalio_common::worker::WorkerDeploymentVersion::builder()
745 .deployment_name("test-deployment".to_string())
746 .build_id("build-success".to_string())
747 .build(),
748 )
749 .use_worker_versioning(true)
750 .build();
751 succeeding_worker
752 .expect_deployment_options()
753 .return_const(success_deployment_options.clone());
754 succeeding_worker
755 .expect_worker_instance_key()
756 .return_const(succeeding_worker_id);
757 succeeding_worker
758 .expect_heartbeat_enabled()
759 .return_const(false);
760 succeeding_worker
761 .expect_worker_task_types()
762 .return_const(WorkerTaskTypes {
763 enable_workflows: true,
764 enable_local_activities: true,
765 enable_remote_activities: true,
766 enable_nexus: true,
767 });
768
769 manager
770 .register(Arc::new(failing_worker), false)
771 .expect("failing worker registration succeeds");
772 manager
773 .register(Arc::new(succeeding_worker), false)
774 .expect("succeeding worker registration succeeds");
775
776 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
777
778 let reservation_deployment_options = reservation
779 .expect("succeeding worker was used after failing worker failed")
780 .deployment_options
781 .unwrap();
782 assert_eq!(
783 reservation_deployment_options, success_deployment_options,
784 "deployment options bubble through from succeeding worker"
785 );
786 }
787
788 #[test]
789 fn reserve_wft_slot_retries_respects_slot_boundary() {
790 let mut manager = ClientWorkerSetImpl::new();
791 let namespace = "retry_namespace".to_string();
792 let task_queue = "retry_queue".to_string();
793
794 let failing_worker_id = Uuid::new_v4();
795 let mut failing_worker = MockClientWorker::new();
796 failing_worker
797 .expect_try_reserve_wft_slot()
798 .times(1)
799 .returning(|| None);
800 failing_worker
801 .expect_namespace()
802 .return_const(namespace.clone());
803 failing_worker
804 .expect_task_queue()
805 .return_const(task_queue.clone());
806 failing_worker.expect_deployment_options().return_const(
807 WorkerDeploymentOptions::new(
808 temporalio_common::worker::WorkerDeploymentVersion::builder()
809 .deployment_name("test-deployment".to_string())
810 .build_id("build-fail".to_string())
811 .build(),
812 )
813 .use_worker_versioning(true)
814 .build(),
815 );
816 failing_worker
817 .expect_worker_instance_key()
818 .return_const(failing_worker_id);
819 failing_worker
820 .expect_heartbeat_enabled()
821 .return_const(false);
822 failing_worker
823 .expect_worker_task_types()
824 .return_const(WorkerTaskTypes {
825 enable_workflows: true,
826 enable_local_activities: true,
827 enable_remote_activities: true,
828 enable_nexus: true,
829 });
830
831 let succeeding_worker_id = Uuid::new_v4();
833 let mut succeeding_worker = MockClientWorker::new();
834 succeeding_worker.expect_try_reserve_wft_slot().times(0);
835 succeeding_worker
836 .expect_namespace()
837 .return_const(namespace.clone());
838 succeeding_worker
839 .expect_task_queue()
840 .return_const("other_task_queue".to_string());
841 succeeding_worker
842 .expect_deployment_options()
843 .return_const(None);
844 succeeding_worker
845 .expect_worker_instance_key()
846 .return_const(succeeding_worker_id);
847 succeeding_worker
848 .expect_heartbeat_enabled()
849 .return_const(false);
850 succeeding_worker
851 .expect_worker_task_types()
852 .return_const(WorkerTaskTypes {
853 enable_workflows: true,
854 enable_local_activities: true,
855 enable_remote_activities: true,
856 enable_nexus: true,
857 });
858
859 manager
860 .register(Arc::new(failing_worker), false)
861 .expect("failing worker registration succeeds");
862 manager
863 .register(Arc::new(succeeding_worker), false)
864 .expect("succeeding worker registration succeeds");
865
866 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
867 assert!(
868 reservation.is_none(),
869 "succeeding_worker should not be picked due to it being on a separate task queue"
870 );
871 }
872
873 #[test]
874 fn registry_keeps_one_provider_per_namespace() {
875 let manager = ClientWorkerSet::new();
876 let mut worker_keys = vec![];
877 let mut successful_registrations = 0;
878
879 for i in 0..10 {
880 let namespace = format!("myId{}", i % 3);
881 let mock_provider =
882 new_mock_provider(namespace, "bar_q".to_string(), false, false, false);
883 let worker_instance_key = mock_provider.worker_instance_key();
884
885 let result = manager.register_worker(Arc::new(mock_provider), false);
886 if let Err(err) = result {
887 assert!(err.to_string().contains(
889 "Registration of multiple workers with overlapping worker task types"
890 ));
891 } else {
892 successful_registrations += 1;
893 worker_keys.push(worker_instance_key);
894 }
895 }
896
897 assert_eq!(successful_registrations, 3);
898 assert_eq!(3, manager.num_providers());
899
900 let count = worker_keys.iter().fold(0, |count, key| {
901 manager.unregister_slot_provider(*key).unwrap();
902 manager.finalize_unregister(*key).unwrap();
903 let result = manager.unregister_slot_provider(*key);
905 assert!(result.is_err());
906 let result = manager.finalize_unregister(*key);
907 assert!(result.is_err());
908 count + 1
909 });
910 assert_eq!(3, count);
911 assert_eq!(0, manager.num_providers());
912 }
913
914 struct MockSharedNamespaceWorker {
915 namespace: String,
916 callbacks: Arc<RwLock<HashMap<Uuid, WorkerCallbacks>>>,
917 worker_control_task_queue_enabled: bool,
918 }
919
920 impl std::fmt::Debug for MockSharedNamespaceWorker {
921 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
922 f.debug_struct("MockSharedNamespaceWorker")
923 .field("namespace", &self.namespace)
924 .field("callbacks_count", &self.callbacks.read().len())
925 .finish()
926 }
927 }
928
929 impl MockSharedNamespaceWorker {
930 fn new(namespace: String) -> Self {
931 Self {
932 namespace,
933 callbacks: Arc::new(RwLock::new(HashMap::new())),
934 worker_control_task_queue_enabled: false,
935 }
936 }
937
938 fn with_worker_control_task_queue_enabled(mut self) -> Self {
939 self.worker_control_task_queue_enabled = true;
940 self
941 }
942 }
943
944 impl SharedNamespaceWorkerTrait for MockSharedNamespaceWorker {
945 fn namespace(&self) -> String {
946 self.namespace.clone()
947 }
948
949 fn register_callback(&self, worker_instance_key: Uuid, callbacks: WorkerCallbacks) {
950 self.callbacks
951 .write()
952 .insert(worker_instance_key, callbacks);
953 }
954
955 fn unregister_callback(
956 &self,
957 worker_instance_key: Uuid,
958 ) -> (Option<WorkerCallbacks>, bool) {
959 let mut callbacks = self.callbacks.write();
960 let callback = callbacks.remove(&worker_instance_key);
961 let is_empty = callbacks.is_empty();
962 (callback, is_empty)
963 }
964
965 fn num_workers(&self) -> usize {
966 self.callbacks.read().len()
967 }
968
969 fn worker_control_task_queue_enabled(&self) -> bool {
970 self.worker_control_task_queue_enabled
971 }
972 }
973
974 #[test]
975 fn worker_control_task_queue_enabled_reflects_shared_worker() {
976 let manager = ClientWorkerSet::new();
977 let namespace = "test_namespace";
978
979 assert!(!manager.worker_control_task_queue_enabled(namespace));
980
981 manager.worker_manager.write().shared_worker.insert(
982 namespace.to_string(),
983 Box::new(
984 MockSharedNamespaceWorker::new(namespace.to_string())
985 .with_worker_control_task_queue_enabled(),
986 ),
987 );
988
989 assert!(manager.worker_control_task_queue_enabled(namespace));
990 assert!(!manager.worker_control_task_queue_enabled("other_namespace"));
991 }
992
993 fn new_mock_provider_with_heartbeat(
994 namespace: String,
995 task_queue: String,
996 heartbeat_enabled: bool,
997 build_id: Option<String>,
998 ) -> MockClientWorker {
999 let mut mock_provider = MockClientWorker::new();
1000 mock_provider
1001 .expect_try_reserve_wft_slot()
1002 .returning(|| Some(new_mock_slot(false)));
1003 mock_provider
1004 .expect_namespace()
1005 .return_const(namespace.clone());
1006 mock_provider.expect_task_queue().return_const(task_queue);
1007 mock_provider
1008 .expect_heartbeat_enabled()
1009 .return_const(heartbeat_enabled);
1010 mock_provider
1011 .expect_worker_instance_key()
1012 .return_const(Uuid::new_v4());
1013 let deployment_name = "test-deployment".to_string();
1014 let build_id_for_closure = build_id.clone();
1015 mock_provider
1016 .expect_deployment_options()
1017 .returning(move || {
1018 build_id_for_closure.as_ref().map(|build_id| {
1019 WorkerDeploymentOptions::new(
1020 temporalio_common::worker::WorkerDeploymentVersion::builder()
1021 .deployment_name(deployment_name.clone())
1022 .build_id(build_id.clone())
1023 .build(),
1024 )
1025 .use_worker_versioning(true)
1026 .build()
1027 })
1028 });
1029
1030 if heartbeat_enabled {
1031 mock_provider
1032 .expect_heartbeat_callback()
1033 .returning(|| Some(Arc::new(|| Some(WorkerHeartbeat::default()))));
1034 mock_provider
1035 .expect_heartbeat_success_callback()
1036 .returning(|| None);
1037 mock_provider
1038 .expect_cancel_activity_callback()
1039 .returning(|| None);
1040
1041 let namespace_clone = namespace.clone();
1042 mock_provider
1043 .expect_new_shared_namespace_worker()
1044 .returning(move || {
1045 Ok(Box::new(MockSharedNamespaceWorker::new(
1046 namespace_clone.clone(),
1047 )))
1048 });
1049 }
1050
1051 mock_provider
1052 .expect_worker_task_types()
1053 .return_const(WorkerTaskTypes {
1054 enable_workflows: true,
1055 enable_local_activities: true,
1056 enable_remote_activities: true,
1057 enable_nexus: true,
1058 });
1059
1060 mock_provider
1061 }
1062
1063 #[test]
1064 fn duplicate_namespace_task_queue_registration_fails() {
1065 let manager = ClientWorkerSet::new();
1066
1067 let worker1 = new_mock_provider_with_heartbeat(
1068 "test_namespace".to_string(),
1069 "test_queue".to_string(),
1070 true,
1071 None,
1072 );
1073
1074 let worker2 = new_mock_provider_with_heartbeat(
1076 "test_namespace".to_string(),
1077 "test_queue".to_string(),
1078 true,
1079 None,
1080 );
1081
1082 manager.register_worker(Arc::new(worker1), false).unwrap();
1083
1084 let result = manager.register_worker(Arc::new(worker2), false);
1086 assert!(result.is_err());
1087 assert!(
1088 result
1089 .unwrap_err()
1090 .to_string()
1091 .contains("Registration of multiple workers with overlapping worker task types")
1092 );
1093
1094 assert_eq!(1, manager.num_providers());
1095 assert_eq!(manager.num_heartbeat_workers(), 1);
1096
1097 let impl_ref = manager.worker_manager.read();
1098 assert_eq!(impl_ref.shared_worker.len(), 1);
1099 assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1100 }
1101
1102 #[test]
1103 fn duplicate_namespace_with_different_build_ids_succeeds() {
1104 let manager = ClientWorkerSet::new();
1105 let namespace = "test_namespace".to_string();
1106 let task_queue = "test_queue".to_string();
1107
1108 let worker1 =
1109 new_mock_provider_with_heartbeat(namespace.clone(), task_queue.clone(), false, None);
1110 let worker1_instance_key = worker1.worker_instance_key();
1111 let worker2 = new_mock_provider_with_heartbeat(
1112 namespace.clone(),
1113 task_queue.clone(),
1114 false,
1115 Some("build-1".to_string()),
1116 );
1117 let worker2_instance_key = worker2.worker_instance_key();
1118 let worker3 =
1119 new_mock_provider_with_heartbeat(namespace.clone(), task_queue.clone(), false, None);
1120 let worker4 = new_mock_provider_with_heartbeat(
1121 namespace.clone(),
1122 task_queue.clone(),
1123 false,
1124 Some("build-1".to_string()),
1125 );
1126
1127 manager.register_worker(Arc::new(worker1), false).unwrap();
1128
1129 manager
1130 .register_worker(Arc::new(worker2), false)
1131 .expect("worker with new build ID should register");
1132 assert_eq!(2, manager.num_providers());
1133
1134 assert!(
1135 manager
1136 .register_worker(Arc::new(worker3), false)
1137 .unwrap_err()
1138 .to_string()
1139 .contains("Registration of multiple workers with overlapping worker task types")
1140 );
1141
1142 assert!(
1143 manager
1144 .register_worker(Arc::new(worker4), false)
1145 .unwrap_err()
1146 .to_string()
1147 .contains("Registration of multiple workers with overlapping worker task types")
1148 );
1149 assert_eq!(2, manager.num_providers());
1150
1151 {
1152 let impl_ref = manager.worker_manager.read();
1153 let slot_key = SlotKey::new(namespace.clone(), task_queue.clone());
1154 let providers = impl_ref
1155 .slot_providers
1156 .get(&slot_key)
1157 .expect("slot providers should exist for namespace/task queue");
1158 assert_eq!(2, providers.len());
1159
1160 assert_eq!(providers[0].worker_id, worker1_instance_key);
1161 assert_eq!(providers[0].build_id, None);
1162 assert_eq!(providers[1].worker_id, worker2_instance_key);
1163 assert_eq!(providers[1].build_id, Some("build-1".to_string()));
1164 }
1165
1166 manager
1167 .unregister_slot_provider(worker2_instance_key)
1168 .unwrap();
1169 manager.finalize_unregister(worker2_instance_key).unwrap();
1170
1171 {
1172 let impl_ref = manager.worker_manager.read();
1173 let slot_key = SlotKey::new(namespace.clone(), task_queue.clone());
1174 let providers = impl_ref
1175 .slot_providers
1176 .get(&slot_key)
1177 .expect("slot providers should exist for namespace/task queue");
1178
1179 assert_eq!(1, providers.len());
1180 assert_eq!(providers[0].worker_id, worker1_instance_key);
1181 assert_eq!(providers[0].build_id, None);
1182 }
1183 }
1184
1185 #[test]
1186 fn multiple_workers_same_namespace_share_heartbeat_manager() {
1187 let manager = ClientWorkerSet::new();
1188
1189 let worker1 = new_mock_provider_with_heartbeat(
1190 "shared_namespace".to_string(),
1191 "queue1".to_string(),
1192 true,
1193 None,
1194 );
1195
1196 let worker2 = new_mock_provider_with_heartbeat(
1198 "shared_namespace".to_string(),
1199 "queue2".to_string(),
1200 true,
1201 None,
1202 );
1203
1204 manager.register_worker(Arc::new(worker1), false).unwrap();
1205 manager.register_worker(Arc::new(worker2), false).unwrap();
1206
1207 assert_eq!(2, manager.num_providers());
1208 assert_eq!(manager.num_heartbeat_workers(), 2);
1209
1210 let impl_ref = manager.worker_manager.read();
1211 assert_eq!(impl_ref.shared_worker.len(), 1);
1212 assert!(impl_ref.shared_worker.contains_key("shared_namespace"));
1213
1214 let shared_worker = impl_ref.shared_worker.get("shared_namespace").unwrap();
1215 assert_eq!(shared_worker.namespace(), "shared_namespace");
1216 }
1217
1218 #[test]
1219 fn different_namespaces_get_separate_heartbeat_managers() {
1220 let manager = ClientWorkerSet::new();
1221 let worker1 = new_mock_provider_with_heartbeat(
1222 "namespace1".to_string(),
1223 "queue1".to_string(),
1224 true,
1225 None,
1226 );
1227 let worker2 = new_mock_provider_with_heartbeat(
1228 "namespace2".to_string(),
1229 "queue1".to_string(),
1230 true,
1231 None,
1232 );
1233
1234 manager.register_worker(Arc::new(worker1), false).unwrap();
1235 manager.register_worker(Arc::new(worker2), false).unwrap();
1236
1237 assert_eq!(2, manager.num_providers());
1238 assert_eq!(manager.num_heartbeat_workers(), 2);
1239
1240 let impl_ref = manager.worker_manager.read();
1241 assert_eq!(impl_ref.num_heartbeat_workers(), 2);
1242 assert!(impl_ref.shared_worker.contains_key("namespace1"));
1243 assert!(impl_ref.shared_worker.contains_key("namespace2"));
1244 }
1245
1246 #[test]
1247 fn unregister_heartbeat_workers_cleans_up_shared_worker_when_last_removed() {
1248 let manager = ClientWorkerSet::new();
1249
1250 let worker1 = new_mock_provider_with_heartbeat(
1252 "test_namespace".to_string(),
1253 "queue1".to_string(),
1254 true,
1255 None,
1256 );
1257 let worker2 = new_mock_provider_with_heartbeat(
1258 "test_namespace".to_string(),
1259 "queue2".to_string(),
1260 true,
1261 None,
1262 );
1263 let worker_instance_key1 = worker1.worker_instance_key();
1264 let worker_instance_key2 = worker2.worker_instance_key();
1265
1266 assert_ne!(worker_instance_key1, worker_instance_key2);
1267
1268 manager.register_worker(Arc::new(worker1), false).unwrap();
1269 manager.register_worker(Arc::new(worker2), false).unwrap();
1270
1271 assert_eq!(2, manager.num_providers());
1273 assert_eq!(manager.num_heartbeat_workers(), 2);
1274
1275 let impl_ref = manager.worker_manager.read();
1276 assert_eq!(impl_ref.shared_worker.len(), 1);
1277 assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1278 assert_eq!(
1279 impl_ref
1280 .shared_worker
1281 .get("test_namespace")
1282 .unwrap()
1283 .num_workers(),
1284 2
1285 );
1286 drop(impl_ref);
1287
1288 manager
1290 .unregister_slot_provider(worker_instance_key1)
1291 .unwrap();
1292 manager.finalize_unregister(worker_instance_key1).unwrap();
1293
1294 assert_eq!(1, manager.num_providers());
1296 assert_eq!(manager.num_heartbeat_workers(), 1);
1297
1298 let impl_ref = manager.worker_manager.read();
1299 assert_eq!(impl_ref.num_heartbeat_workers(), 1); assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1301 assert_eq!(
1302 impl_ref
1303 .shared_worker
1304 .get("test_namespace")
1305 .unwrap()
1306 .num_workers(),
1307 1
1308 );
1309 drop(impl_ref);
1310
1311 manager
1313 .unregister_slot_provider(worker_instance_key2)
1314 .unwrap();
1315 manager.finalize_unregister(worker_instance_key2).unwrap();
1316
1317 assert_eq!(0, manager.num_providers());
1319 assert_eq!(manager.num_heartbeat_workers(), 0);
1320
1321 let impl_ref = manager.worker_manager.read();
1322 assert_eq!(impl_ref.shared_worker.len(), 0); assert!(!impl_ref.shared_worker.contains_key("test_namespace"));
1324 }
1325
1326 #[test]
1327 fn workflow_and_activity_only_workers_coexist() {
1328 let manager = ClientWorkerSet::new();
1329 let namespace = "test_namespace".to_string();
1330 let task_queue = "test_queue".to_string();
1331
1332 let mut workflow_nexus_worker = MockClientWorker::new();
1333 workflow_nexus_worker
1334 .expect_namespace()
1335 .return_const(namespace.clone());
1336 workflow_nexus_worker
1337 .expect_task_queue()
1338 .return_const(task_queue.clone());
1339 workflow_nexus_worker
1340 .expect_deployment_options()
1341 .return_const(None);
1342 workflow_nexus_worker
1343 .expect_worker_instance_key()
1344 .return_const(Uuid::new_v4());
1345 workflow_nexus_worker
1346 .expect_heartbeat_enabled()
1347 .return_const(false);
1348 workflow_nexus_worker
1349 .expect_worker_task_types()
1350 .return_const(WorkerTaskTypes {
1351 enable_workflows: true,
1352 enable_local_activities: false,
1353 enable_remote_activities: false,
1354 enable_nexus: true,
1355 });
1356
1357 let mut activity_worker = MockClientWorker::new();
1358 activity_worker
1359 .expect_namespace()
1360 .return_const(namespace.clone());
1361 activity_worker
1362 .expect_task_queue()
1363 .return_const(task_queue.clone());
1364 activity_worker
1365 .expect_deployment_options()
1366 .return_const(None);
1367 activity_worker
1368 .expect_worker_instance_key()
1369 .return_const(Uuid::new_v4());
1370 activity_worker
1371 .expect_heartbeat_enabled()
1372 .return_const(false);
1373 activity_worker
1374 .expect_worker_task_types()
1375 .return_const(WorkerTaskTypes {
1376 enable_workflows: false,
1377 enable_local_activities: false,
1378 enable_remote_activities: true,
1379 enable_nexus: false,
1380 });
1381 activity_worker.expect_try_reserve_wft_slot().times(0); manager
1384 .register_worker(Arc::new(workflow_nexus_worker), false)
1385 .expect("workflow-nexus worker should register");
1386 manager
1387 .register_worker(Arc::new(activity_worker), false)
1388 .expect("activity-only worker should register");
1389
1390 assert_eq!(2, manager.num_providers());
1391 }
1392
1393 #[test]
1394 fn overlapping_capabilities_rejected() {
1395 let manager = ClientWorkerSet::new();
1396 let namespace = "test_namespace".to_string();
1397 let task_queue = "test_queue".to_string();
1398
1399 let mut worker1 = MockClientWorker::new();
1401 worker1.expect_namespace().return_const(namespace.clone());
1402 worker1.expect_task_queue().return_const(task_queue.clone());
1403 worker1.expect_deployment_options().return_const(None);
1404 worker1
1405 .expect_worker_instance_key()
1406 .return_const(Uuid::new_v4());
1407 worker1.expect_heartbeat_enabled().return_const(false);
1408 worker1
1409 .expect_worker_task_types()
1410 .return_const(WorkerTaskTypes {
1411 enable_workflows: true,
1412 enable_local_activities: true,
1413 enable_remote_activities: true,
1414 enable_nexus: false,
1415 });
1416
1417 let mut worker2 = MockClientWorker::new();
1419 worker2.expect_namespace().return_const(namespace.clone());
1420 worker2.expect_task_queue().return_const(task_queue.clone());
1421 worker2.expect_deployment_options().return_const(None);
1422 worker2
1423 .expect_worker_instance_key()
1424 .return_const(Uuid::new_v4());
1425 worker2.expect_heartbeat_enabled().return_const(false);
1426 worker2
1427 .expect_worker_task_types()
1428 .return_const(WorkerTaskTypes {
1429 enable_workflows: true,
1430 enable_local_activities: true,
1431 enable_remote_activities: true,
1432 enable_nexus: false,
1433 });
1434
1435 manager
1436 .register_worker(Arc::new(worker1), false)
1437 .expect("first worker should register");
1438
1439 let result = manager.register_worker(Arc::new(worker2), false);
1440 assert!(result.is_err());
1441 assert!(
1442 result
1443 .unwrap_err()
1444 .to_string()
1445 .contains("overlapping worker task types")
1446 );
1447
1448 let mut worker3 = MockClientWorker::new();
1450 worker3.expect_namespace().return_const(namespace.clone());
1451 worker3.expect_task_queue().return_const(task_queue.clone());
1452 worker3.expect_deployment_options().return_const(None);
1453 worker3
1454 .expect_worker_instance_key()
1455 .return_const(Uuid::new_v4());
1456 worker3.expect_heartbeat_enabled().return_const(false);
1457 worker3
1458 .expect_worker_task_types()
1459 .return_const(WorkerTaskTypes {
1460 enable_workflows: false,
1461 enable_local_activities: false,
1462 enable_remote_activities: true,
1463 enable_nexus: false,
1464 });
1465
1466 let result = manager.register_worker(Arc::new(worker3), false);
1467 assert!(result.is_err());
1468 assert!(
1469 result
1470 .unwrap_err()
1471 .to_string()
1472 .contains("overlapping worker task types")
1473 );
1474 }
1475
1476 #[test]
1477 fn wft_slot_reservation_ignores_non_workflow_workers() {
1478 let mut manager_impl = ClientWorkerSetImpl::new();
1479 let namespace = "test_namespace".to_string();
1480 let task_queue = "test_queue".to_string();
1481
1482 let mut activity_worker = MockClientWorker::new();
1483 activity_worker
1484 .expect_namespace()
1485 .return_const(namespace.clone());
1486 activity_worker
1487 .expect_task_queue()
1488 .return_const(task_queue.clone());
1489 activity_worker
1490 .expect_deployment_options()
1491 .return_const(None);
1492 activity_worker
1493 .expect_worker_instance_key()
1494 .return_const(Uuid::new_v4());
1495 activity_worker
1496 .expect_heartbeat_enabled()
1497 .return_const(false);
1498 activity_worker
1499 .expect_worker_task_types()
1500 .return_const(WorkerTaskTypes {
1501 enable_workflows: false,
1502 enable_local_activities: false,
1503 enable_remote_activities: true,
1504 enable_nexus: false,
1505 });
1506
1507 let mut nexus_worker = MockClientWorker::new();
1508 nexus_worker
1509 .expect_namespace()
1510 .return_const(namespace.clone());
1511 nexus_worker
1512 .expect_task_queue()
1513 .return_const(task_queue.clone());
1514 nexus_worker.expect_deployment_options().return_const(None);
1515 nexus_worker
1516 .expect_worker_instance_key()
1517 .return_const(Uuid::new_v4());
1518 nexus_worker.expect_heartbeat_enabled().return_const(false);
1519 nexus_worker
1520 .expect_worker_task_types()
1521 .return_const(WorkerTaskTypes {
1522 enable_workflows: false,
1523 enable_local_activities: false,
1524 enable_remote_activities: false,
1525 enable_nexus: true,
1526 });
1527
1528 manager_impl
1529 .register(Arc::new(activity_worker), false)
1530 .expect("activity worker should register");
1531 manager_impl
1532 .register(Arc::new(nexus_worker), false)
1533 .expect("nexus worker should register");
1534
1535 let reservation = manager_impl.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1536 assert!(
1537 reservation.is_none(),
1538 "should not find workflow workers when only activity/nexus workers registered"
1539 );
1540
1541 let mut workflow_worker = MockClientWorker::new();
1543 workflow_worker
1544 .expect_namespace()
1545 .return_const(namespace.clone());
1546 workflow_worker
1547 .expect_task_queue()
1548 .return_const(task_queue.clone());
1549 workflow_worker
1550 .expect_deployment_options()
1551 .return_const(None);
1552 workflow_worker
1553 .expect_worker_instance_key()
1554 .return_const(Uuid::new_v4());
1555 workflow_worker
1556 .expect_heartbeat_enabled()
1557 .return_const(false);
1558 workflow_worker
1559 .expect_worker_task_types()
1560 .return_const(WorkerTaskTypes {
1561 enable_workflows: true,
1562 enable_local_activities: true,
1563 enable_remote_activities: false,
1564 enable_nexus: false,
1565 });
1566 workflow_worker
1567 .expect_try_reserve_wft_slot()
1568 .times(1)
1569 .returning(|| Some(new_mock_slot(false)));
1570
1571 manager_impl
1572 .register(Arc::new(workflow_worker), false)
1573 .expect("workflow worker should register");
1574
1575 let reservation = manager_impl.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1576 assert!(
1577 reservation.is_some(),
1578 "should find workflow worker after it's registered"
1579 );
1580 }
1581
1582 #[test]
1583 fn worker_invalid_type_config_rejected() {
1584 let manager = ClientWorkerSet::new();
1585
1586 let mut worker = MockClientWorker::new();
1588 worker
1589 .expect_namespace()
1590 .return_const("test_namespace".to_string());
1591 worker
1592 .expect_task_queue()
1593 .return_const("test_queue".to_string());
1594 worker.expect_deployment_options().return_const(None);
1595 worker
1596 .expect_worker_instance_key()
1597 .return_const(Uuid::new_v4());
1598 worker.expect_heartbeat_enabled().return_const(false);
1599 worker
1600 .expect_worker_task_types()
1601 .return_const(WorkerTaskTypes {
1602 enable_workflows: false,
1603 enable_local_activities: false,
1604 enable_remote_activities: false,
1605 enable_nexus: false,
1606 });
1607
1608 let result = manager.register_worker(Arc::new(worker), false);
1609 assert!(result.is_err());
1610 assert!(
1611 result
1612 .unwrap_err()
1613 .to_string()
1614 .contains("must have at least one capability enabled")
1615 );
1616
1617 let mut worker = MockClientWorker::new();
1619 worker
1620 .expect_namespace()
1621 .return_const("test_namespace".to_string());
1622 worker
1623 .expect_task_queue()
1624 .return_const("test_queue".to_string());
1625 worker.expect_deployment_options().return_const(None);
1626 worker
1627 .expect_worker_instance_key()
1628 .return_const(Uuid::new_v4());
1629 worker.expect_heartbeat_enabled().return_const(false);
1630 worker
1631 .expect_worker_task_types()
1632 .return_const(WorkerTaskTypes {
1633 enable_workflows: false,
1634 enable_local_activities: true,
1635 enable_remote_activities: true,
1636 enable_nexus: false,
1637 });
1638
1639 let result = manager.register_worker(Arc::new(worker), false);
1640 assert!(result.is_err());
1641 assert_eq!(
1642 result.unwrap_err().to_string(),
1643 "Local activities cannot be enabled without workflows".to_string()
1644 );
1645 }
1646
1647 #[test]
1648 fn unregister_with_multiple_workers() {
1649 let manager = ClientWorkerSet::new();
1650 let namespace = "test_namespace".to_string();
1651 let task_queue = "test_queue".to_string();
1652
1653 let mut workflow_worker = MockClientWorker::new();
1655 workflow_worker
1656 .expect_namespace()
1657 .return_const(namespace.clone());
1658 workflow_worker
1659 .expect_task_queue()
1660 .return_const(task_queue.clone());
1661 workflow_worker
1662 .expect_deployment_options()
1663 .return_const(None);
1664 let wf_worker_key = Uuid::new_v4();
1665 workflow_worker
1666 .expect_worker_instance_key()
1667 .return_const(wf_worker_key);
1668 workflow_worker
1669 .expect_heartbeat_enabled()
1670 .return_const(false);
1671 workflow_worker
1672 .expect_worker_task_types()
1673 .return_const(WorkerTaskTypes {
1674 enable_workflows: true,
1675 enable_local_activities: true,
1676 enable_remote_activities: false,
1677 enable_nexus: false,
1678 });
1679 workflow_worker
1680 .expect_try_reserve_wft_slot()
1681 .returning(|| Some(new_mock_slot(false)));
1682
1683 let mut activity_worker = MockClientWorker::new();
1685 activity_worker
1686 .expect_namespace()
1687 .return_const(namespace.clone());
1688 activity_worker
1689 .expect_task_queue()
1690 .return_const(task_queue.clone());
1691 activity_worker
1692 .expect_deployment_options()
1693 .return_const(None);
1694 let act_worker_key = Uuid::new_v4();
1695 activity_worker
1696 .expect_worker_instance_key()
1697 .return_const(act_worker_key);
1698 activity_worker
1699 .expect_heartbeat_enabled()
1700 .return_const(false);
1701 activity_worker
1702 .expect_worker_task_types()
1703 .return_const(WorkerTaskTypes {
1704 enable_workflows: false,
1705 enable_local_activities: false,
1706 enable_remote_activities: true,
1707 enable_nexus: false,
1708 });
1709
1710 manager
1711 .register_worker(Arc::new(workflow_worker), false)
1712 .expect("workflow worker should register");
1713 manager
1714 .register_worker(Arc::new(activity_worker), false)
1715 .expect("activity worker should register");
1716
1717 assert_eq!(2, manager.num_providers());
1718
1719 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1720 assert!(
1721 reservation.is_some(),
1722 "should be able to reserve slot from workflow worker"
1723 );
1724
1725 manager
1726 .unregister_slot_provider(wf_worker_key)
1727 .expect("should unregister slot provider for workflow worker");
1728 manager
1729 .finalize_unregister(wf_worker_key)
1730 .expect("should finalize unregister for workflow worker");
1731
1732 assert_eq!(1, manager.num_providers());
1734
1735 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1736 assert!(
1737 reservation.is_none(),
1738 "should not find workflow worker after unregistration"
1739 );
1740
1741 manager
1742 .unregister_slot_provider(act_worker_key)
1743 .expect("should unregister slot provider for activity worker");
1744 manager
1745 .finalize_unregister(act_worker_key)
1746 .expect("should finalize unregister for activity worker");
1747
1748 assert_eq!(0, manager.num_providers());
1749 }
1750
1751 #[test]
1752 fn worker_unregister_order() {
1753 let manager = ClientWorkerSet::new();
1754 let worker = new_mock_provider_with_heartbeat(
1755 "namespace1".to_string(),
1756 "queue1".to_string(),
1757 true,
1758 None,
1759 );
1760 let worker_instance_key = worker.worker_instance_key();
1761 manager.register_worker(Arc::new(worker), false).unwrap();
1762
1763 let res = manager.finalize_unregister(worker_instance_key);
1764 assert!(res.is_err());
1765 let err_string = res.err().map(|e| e.to_string()).unwrap();
1766 assert!(err_string.contains("Worker still in slot_providers during finalize"));
1767
1768 manager
1771 .unregister_slot_provider(worker_instance_key)
1772 .unwrap();
1773 manager.finalize_unregister(worker_instance_key).unwrap();
1774 }
1775}