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