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 {
234 heartbeat: heartbeat_callback,
235 heartbeat_success: worker.heartbeat_success_callback(),
236 cancel_activity: 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
513pub struct WorkerCallbacks {
515 pub heartbeat: HeartbeatCallback,
517 pub heartbeat_success: Option<HeartbeatSuccessCallback>,
519 pub cancel_activity: Option<CancelActivityCallback>,
521}
522
523#[cfg_attr(test, mockall::automock)]
526pub trait ClientWorker: Send + Sync {
527 fn namespace(&self) -> &str;
529
530 fn task_queue(&self) -> &str;
532
533 fn try_reserve_wft_slot(&self) -> Option<Box<dyn Slot + Send>>;
539
540 fn deployment_options(&self) -> Option<WorkerDeploymentOptions>;
542
543 fn worker_instance_key(&self) -> Uuid;
546
547 fn heartbeat_enabled(&self) -> bool;
549
550 fn heartbeat_callback(&self) -> Option<HeartbeatCallback>;
552
553 fn heartbeat_success_callback(&self) -> Option<HeartbeatSuccessCallback> {
555 None
556 }
557
558 fn cancel_activity_callback(&self) -> Option<CancelActivityCallback>;
560
561 fn new_shared_namespace_worker(
563 &self,
564 ) -> Result<Box<dyn SharedNamespaceWorkerTrait + Send + Sync>, anyhow::Error>;
565
566 fn worker_task_types(&self) -> WorkerTaskTypes;
568}
569
570#[cfg(test)]
571mod tests {
572 use super::*;
573 use std::sync::atomic::{AtomicUsize, Ordering};
574
575 #[tokio::test]
576 async fn namespace_description_source_resolves_once() {
577 let source = NamespaceDescriptionSource::unresolved();
578 let calls = AtomicUsize::new(0);
579
580 let first = source.resolve(|| async {
581 calls.fetch_add(1, Ordering::Relaxed);
582 tokio::task::yield_now().await;
583 Ok(DescribeNamespaceResponse::default())
584 });
585 let second = source.resolve(|| async {
586 calls.fetch_add(1, Ordering::Relaxed);
587 Ok(DescribeNamespaceResponse::default())
588 });
589
590 let (first, second) = tokio::join!(first, second);
591 first.unwrap();
592 second.unwrap();
593 assert_eq!(calls.load(Ordering::Relaxed), 1);
594 }
595
596 #[tokio::test]
597 async fn namespace_description_source_resolves_unimplemented_as_default() {
598 let source = NamespaceDescriptionSource::unresolved();
599
600 let description = source
601 .resolve(|| async { Err(tonic::Status::unimplemented("unsupported")) })
602 .await
603 .unwrap();
604
605 assert_eq!(description, &DescribeNamespaceResponse::default());
606 }
607
608 #[test]
609 fn namespace_description_sources_are_scoped_by_namespace() {
610 let workers = ClientWorkerSet::new();
611 let first = workers.namespace_description_source("first");
612
613 assert!(Arc::ptr_eq(
614 &first,
615 &workers.namespace_description_source("first")
616 ));
617 assert!(!Arc::ptr_eq(
618 &first,
619 &workers.namespace_description_source("second")
620 ));
621 }
622
623 fn new_mock_slot(with_error: bool) -> Box<MockSlot> {
624 let mut mock_slot = MockSlot::new();
625 if with_error {
626 mock_slot
627 .expect_schedule_wft()
628 .returning(|_| Err(anyhow::anyhow!("Changed my mind")));
629 } else {
630 mock_slot.expect_schedule_wft().returning(|_| Ok(()));
631 }
632 Box::new(mock_slot)
633 }
634
635 fn new_mock_provider(
636 namespace: String,
637 task_queue: String,
638 with_error: bool,
639 no_slots: bool,
640 heartbeat_enabled: bool,
641 ) -> MockClientWorker {
642 let mut mock_provider = MockClientWorker::new();
643 mock_provider
644 .expect_try_reserve_wft_slot()
645 .returning(move || {
646 if no_slots {
647 None
648 } else {
649 Some(new_mock_slot(with_error))
650 }
651 });
652 mock_provider.expect_namespace().return_const(namespace);
653 mock_provider.expect_task_queue().return_const(task_queue);
654 mock_provider.expect_deployment_options().return_const(None);
655 mock_provider
656 .expect_heartbeat_enabled()
657 .return_const(heartbeat_enabled);
658 mock_provider
659 .expect_worker_instance_key()
660 .return_const(Uuid::new_v4());
661 mock_provider
662 .expect_worker_task_types()
663 .return_const(WorkerTaskTypes {
664 enable_workflows: true,
665 enable_local_activities: true,
666 enable_remote_activities: true,
667 enable_nexus: true,
668 });
669 mock_provider
670 }
671
672 #[test]
673 fn reserve_wft_slot_retries_another_worker_when_first_has_no_slot() {
674 let mut manager = ClientWorkerSetImpl::new();
675 let namespace = "retry_namespace".to_string();
676 let task_queue = "retry_queue".to_string();
677
678 let failing_worker_id = Uuid::new_v4();
679 let mut failing_worker = MockClientWorker::new();
680 failing_worker
681 .expect_try_reserve_wft_slot()
682 .times(1)
683 .returning(|| None);
684 failing_worker
685 .expect_namespace()
686 .return_const(namespace.clone());
687 failing_worker
688 .expect_task_queue()
689 .return_const(task_queue.clone());
690 failing_worker.expect_deployment_options().return_const(
691 WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion {
692 deployment_name: "test-deployment".to_string(),
693 build_id: "build-fail".to_string(),
694 })
695 .use_worker_versioning(true)
696 .build(),
697 );
698 failing_worker
699 .expect_worker_instance_key()
700 .return_const(failing_worker_id);
701 failing_worker
702 .expect_heartbeat_enabled()
703 .return_const(false);
704 failing_worker
705 .expect_worker_task_types()
706 .return_const(WorkerTaskTypes {
707 enable_workflows: true,
708 enable_local_activities: true,
709 enable_remote_activities: true,
710 enable_nexus: true,
711 });
712
713 let succeeding_worker_id = Uuid::new_v4();
714 let mut succeeding_worker = MockClientWorker::new();
715 succeeding_worker
716 .expect_try_reserve_wft_slot()
717 .times(1)
718 .returning(|| Some(new_mock_slot(false)));
719 succeeding_worker
720 .expect_namespace()
721 .return_const(namespace.clone());
722 succeeding_worker
723 .expect_task_queue()
724 .return_const(task_queue.clone());
725 let success_deployment_options =
726 WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion {
727 deployment_name: "test-deployment".to_string(),
728 build_id: "build-success".to_string(),
729 })
730 .use_worker_versioning(true)
731 .build();
732 succeeding_worker
733 .expect_deployment_options()
734 .return_const(success_deployment_options.clone());
735 succeeding_worker
736 .expect_worker_instance_key()
737 .return_const(succeeding_worker_id);
738 succeeding_worker
739 .expect_heartbeat_enabled()
740 .return_const(false);
741 succeeding_worker
742 .expect_worker_task_types()
743 .return_const(WorkerTaskTypes {
744 enable_workflows: true,
745 enable_local_activities: true,
746 enable_remote_activities: true,
747 enable_nexus: true,
748 });
749
750 manager
751 .register(Arc::new(failing_worker), false)
752 .expect("failing worker registration succeeds");
753 manager
754 .register(Arc::new(succeeding_worker), false)
755 .expect("succeeding worker registration succeeds");
756
757 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
758
759 let reservation_deployment_options = reservation
760 .expect("succeeding worker was used after failing worker failed")
761 .deployment_options
762 .unwrap();
763 assert_eq!(
764 reservation_deployment_options, success_deployment_options,
765 "deployment options bubble through from succeeding worker"
766 );
767 }
768
769 #[test]
770 fn reserve_wft_slot_retries_respects_slot_boundary() {
771 let mut manager = ClientWorkerSetImpl::new();
772 let namespace = "retry_namespace".to_string();
773 let task_queue = "retry_queue".to_string();
774
775 let failing_worker_id = Uuid::new_v4();
776 let mut failing_worker = MockClientWorker::new();
777 failing_worker
778 .expect_try_reserve_wft_slot()
779 .times(1)
780 .returning(|| None);
781 failing_worker
782 .expect_namespace()
783 .return_const(namespace.clone());
784 failing_worker
785 .expect_task_queue()
786 .return_const(task_queue.clone());
787 failing_worker.expect_deployment_options().return_const(
788 WorkerDeploymentOptions::new(temporalio_common::worker::WorkerDeploymentVersion {
789 deployment_name: "test-deployment".to_string(),
790 build_id: "build-fail".to_string(),
791 })
792 .use_worker_versioning(true)
793 .build(),
794 );
795 failing_worker
796 .expect_worker_instance_key()
797 .return_const(failing_worker_id);
798 failing_worker
799 .expect_heartbeat_enabled()
800 .return_const(false);
801 failing_worker
802 .expect_worker_task_types()
803 .return_const(WorkerTaskTypes {
804 enable_workflows: true,
805 enable_local_activities: true,
806 enable_remote_activities: true,
807 enable_nexus: true,
808 });
809
810 let succeeding_worker_id = Uuid::new_v4();
812 let mut succeeding_worker = MockClientWorker::new();
813 succeeding_worker.expect_try_reserve_wft_slot().times(0);
814 succeeding_worker
815 .expect_namespace()
816 .return_const(namespace.clone());
817 succeeding_worker
818 .expect_task_queue()
819 .return_const("other_task_queue".to_string());
820 succeeding_worker
821 .expect_deployment_options()
822 .return_const(None);
823 succeeding_worker
824 .expect_worker_instance_key()
825 .return_const(succeeding_worker_id);
826 succeeding_worker
827 .expect_heartbeat_enabled()
828 .return_const(false);
829 succeeding_worker
830 .expect_worker_task_types()
831 .return_const(WorkerTaskTypes {
832 enable_workflows: true,
833 enable_local_activities: true,
834 enable_remote_activities: true,
835 enable_nexus: true,
836 });
837
838 manager
839 .register(Arc::new(failing_worker), false)
840 .expect("failing worker registration succeeds");
841 manager
842 .register(Arc::new(succeeding_worker), false)
843 .expect("succeeding worker registration succeeds");
844
845 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
846 assert!(
847 reservation.is_none(),
848 "succeeding_worker should not be picked due to it being on a separate task queue"
849 );
850 }
851
852 #[test]
853 fn registry_keeps_one_provider_per_namespace() {
854 let manager = ClientWorkerSet::new();
855 let mut worker_keys = vec![];
856 let mut successful_registrations = 0;
857
858 for i in 0..10 {
859 let namespace = format!("myId{}", i % 3);
860 let mock_provider =
861 new_mock_provider(namespace, "bar_q".to_string(), false, false, false);
862 let worker_instance_key = mock_provider.worker_instance_key();
863
864 let result = manager.register_worker(Arc::new(mock_provider), false);
865 if let Err(err) = result {
866 assert!(err.to_string().contains(
868 "Registration of multiple workers with overlapping worker task types"
869 ));
870 } else {
871 successful_registrations += 1;
872 worker_keys.push(worker_instance_key);
873 }
874 }
875
876 assert_eq!(successful_registrations, 3);
877 assert_eq!(3, manager.num_providers());
878
879 let count = worker_keys.iter().fold(0, |count, key| {
880 manager.unregister_slot_provider(*key).unwrap();
881 manager.finalize_unregister(*key).unwrap();
882 let result = manager.unregister_slot_provider(*key);
884 assert!(result.is_err());
885 let result = manager.finalize_unregister(*key);
886 assert!(result.is_err());
887 count + 1
888 });
889 assert_eq!(3, count);
890 assert_eq!(0, manager.num_providers());
891 }
892
893 struct MockSharedNamespaceWorker {
894 namespace: String,
895 callbacks: Arc<RwLock<HashMap<Uuid, WorkerCallbacks>>>,
896 worker_control_task_queue_enabled: bool,
897 }
898
899 impl std::fmt::Debug for MockSharedNamespaceWorker {
900 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
901 f.debug_struct("MockSharedNamespaceWorker")
902 .field("namespace", &self.namespace)
903 .field("callbacks_count", &self.callbacks.read().len())
904 .finish()
905 }
906 }
907
908 impl MockSharedNamespaceWorker {
909 fn new(namespace: String) -> Self {
910 Self {
911 namespace,
912 callbacks: Arc::new(RwLock::new(HashMap::new())),
913 worker_control_task_queue_enabled: false,
914 }
915 }
916
917 fn with_worker_control_task_queue_enabled(mut self) -> Self {
918 self.worker_control_task_queue_enabled = true;
919 self
920 }
921 }
922
923 impl SharedNamespaceWorkerTrait for MockSharedNamespaceWorker {
924 fn namespace(&self) -> String {
925 self.namespace.clone()
926 }
927
928 fn register_callback(&self, worker_instance_key: Uuid, callbacks: WorkerCallbacks) {
929 self.callbacks
930 .write()
931 .insert(worker_instance_key, callbacks);
932 }
933
934 fn unregister_callback(
935 &self,
936 worker_instance_key: Uuid,
937 ) -> (Option<WorkerCallbacks>, bool) {
938 let mut callbacks = self.callbacks.write();
939 let callback = callbacks.remove(&worker_instance_key);
940 let is_empty = callbacks.is_empty();
941 (callback, is_empty)
942 }
943
944 fn num_workers(&self) -> usize {
945 self.callbacks.read().len()
946 }
947
948 fn worker_control_task_queue_enabled(&self) -> bool {
949 self.worker_control_task_queue_enabled
950 }
951 }
952
953 #[test]
954 fn worker_control_task_queue_enabled_reflects_shared_worker() {
955 let manager = ClientWorkerSet::new();
956 let namespace = "test_namespace";
957
958 assert!(!manager.worker_control_task_queue_enabled(namespace));
959
960 manager.worker_manager.write().shared_worker.insert(
961 namespace.to_string(),
962 Box::new(
963 MockSharedNamespaceWorker::new(namespace.to_string())
964 .with_worker_control_task_queue_enabled(),
965 ),
966 );
967
968 assert!(manager.worker_control_task_queue_enabled(namespace));
969 assert!(!manager.worker_control_task_queue_enabled("other_namespace"));
970 }
971
972 fn new_mock_provider_with_heartbeat(
973 namespace: String,
974 task_queue: String,
975 heartbeat_enabled: bool,
976 build_id: Option<String>,
977 ) -> MockClientWorker {
978 let mut mock_provider = MockClientWorker::new();
979 mock_provider
980 .expect_try_reserve_wft_slot()
981 .returning(|| Some(new_mock_slot(false)));
982 mock_provider
983 .expect_namespace()
984 .return_const(namespace.clone());
985 mock_provider.expect_task_queue().return_const(task_queue);
986 mock_provider
987 .expect_heartbeat_enabled()
988 .return_const(heartbeat_enabled);
989 mock_provider
990 .expect_worker_instance_key()
991 .return_const(Uuid::new_v4());
992 let deployment_name = "test-deployment".to_string();
993 let build_id_for_closure = build_id.clone();
994 mock_provider
995 .expect_deployment_options()
996 .returning(move || {
997 build_id_for_closure.as_ref().map(|build_id| {
998 WorkerDeploymentOptions::new(
999 temporalio_common::worker::WorkerDeploymentVersion {
1000 deployment_name: deployment_name.clone(),
1001 build_id: build_id.clone(),
1002 },
1003 )
1004 .use_worker_versioning(true)
1005 .build()
1006 })
1007 });
1008
1009 if heartbeat_enabled {
1010 mock_provider
1011 .expect_heartbeat_callback()
1012 .returning(|| Some(Arc::new(|| Some(WorkerHeartbeat::default()))));
1013 mock_provider
1014 .expect_heartbeat_success_callback()
1015 .returning(|| None);
1016 mock_provider
1017 .expect_cancel_activity_callback()
1018 .returning(|| None);
1019
1020 let namespace_clone = namespace.clone();
1021 mock_provider
1022 .expect_new_shared_namespace_worker()
1023 .returning(move || {
1024 Ok(Box::new(MockSharedNamespaceWorker::new(
1025 namespace_clone.clone(),
1026 )))
1027 });
1028 }
1029
1030 mock_provider
1031 .expect_worker_task_types()
1032 .return_const(WorkerTaskTypes {
1033 enable_workflows: true,
1034 enable_local_activities: true,
1035 enable_remote_activities: true,
1036 enable_nexus: true,
1037 });
1038
1039 mock_provider
1040 }
1041
1042 #[test]
1043 fn duplicate_namespace_task_queue_registration_fails() {
1044 let manager = ClientWorkerSet::new();
1045
1046 let worker1 = new_mock_provider_with_heartbeat(
1047 "test_namespace".to_string(),
1048 "test_queue".to_string(),
1049 true,
1050 None,
1051 );
1052
1053 let worker2 = new_mock_provider_with_heartbeat(
1055 "test_namespace".to_string(),
1056 "test_queue".to_string(),
1057 true,
1058 None,
1059 );
1060
1061 manager.register_worker(Arc::new(worker1), false).unwrap();
1062
1063 let result = manager.register_worker(Arc::new(worker2), false);
1065 assert!(result.is_err());
1066 assert!(
1067 result
1068 .unwrap_err()
1069 .to_string()
1070 .contains("Registration of multiple workers with overlapping worker task types")
1071 );
1072
1073 assert_eq!(1, manager.num_providers());
1074 assert_eq!(manager.num_heartbeat_workers(), 1);
1075
1076 let impl_ref = manager.worker_manager.read();
1077 assert_eq!(impl_ref.shared_worker.len(), 1);
1078 assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1079 }
1080
1081 #[test]
1082 fn duplicate_namespace_with_different_build_ids_succeeds() {
1083 let manager = ClientWorkerSet::new();
1084 let namespace = "test_namespace".to_string();
1085 let task_queue = "test_queue".to_string();
1086
1087 let worker1 =
1088 new_mock_provider_with_heartbeat(namespace.clone(), task_queue.clone(), false, None);
1089 let worker1_instance_key = worker1.worker_instance_key();
1090 let worker2 = new_mock_provider_with_heartbeat(
1091 namespace.clone(),
1092 task_queue.clone(),
1093 false,
1094 Some("build-1".to_string()),
1095 );
1096 let worker2_instance_key = worker2.worker_instance_key();
1097 let worker3 =
1098 new_mock_provider_with_heartbeat(namespace.clone(), task_queue.clone(), false, None);
1099 let worker4 = new_mock_provider_with_heartbeat(
1100 namespace.clone(),
1101 task_queue.clone(),
1102 false,
1103 Some("build-1".to_string()),
1104 );
1105
1106 manager.register_worker(Arc::new(worker1), false).unwrap();
1107
1108 manager
1109 .register_worker(Arc::new(worker2), false)
1110 .expect("worker with new build ID should register");
1111 assert_eq!(2, manager.num_providers());
1112
1113 assert!(
1114 manager
1115 .register_worker(Arc::new(worker3), false)
1116 .unwrap_err()
1117 .to_string()
1118 .contains("Registration of multiple workers with overlapping worker task types")
1119 );
1120
1121 assert!(
1122 manager
1123 .register_worker(Arc::new(worker4), false)
1124 .unwrap_err()
1125 .to_string()
1126 .contains("Registration of multiple workers with overlapping worker task types")
1127 );
1128 assert_eq!(2, manager.num_providers());
1129
1130 {
1131 let impl_ref = manager.worker_manager.read();
1132 let slot_key = SlotKey::new(namespace.clone(), task_queue.clone());
1133 let providers = impl_ref
1134 .slot_providers
1135 .get(&slot_key)
1136 .expect("slot providers should exist for namespace/task queue");
1137 assert_eq!(2, providers.len());
1138
1139 assert_eq!(providers[0].worker_id, worker1_instance_key);
1140 assert_eq!(providers[0].build_id, None);
1141 assert_eq!(providers[1].worker_id, worker2_instance_key);
1142 assert_eq!(providers[1].build_id, Some("build-1".to_string()));
1143 }
1144
1145 manager
1146 .unregister_slot_provider(worker2_instance_key)
1147 .unwrap();
1148 manager.finalize_unregister(worker2_instance_key).unwrap();
1149
1150 {
1151 let impl_ref = manager.worker_manager.read();
1152 let slot_key = SlotKey::new(namespace.clone(), task_queue.clone());
1153 let providers = impl_ref
1154 .slot_providers
1155 .get(&slot_key)
1156 .expect("slot providers should exist for namespace/task queue");
1157
1158 assert_eq!(1, providers.len());
1159 assert_eq!(providers[0].worker_id, worker1_instance_key);
1160 assert_eq!(providers[0].build_id, None);
1161 }
1162 }
1163
1164 #[test]
1165 fn multiple_workers_same_namespace_share_heartbeat_manager() {
1166 let manager = ClientWorkerSet::new();
1167
1168 let worker1 = new_mock_provider_with_heartbeat(
1169 "shared_namespace".to_string(),
1170 "queue1".to_string(),
1171 true,
1172 None,
1173 );
1174
1175 let worker2 = new_mock_provider_with_heartbeat(
1177 "shared_namespace".to_string(),
1178 "queue2".to_string(),
1179 true,
1180 None,
1181 );
1182
1183 manager.register_worker(Arc::new(worker1), false).unwrap();
1184 manager.register_worker(Arc::new(worker2), false).unwrap();
1185
1186 assert_eq!(2, manager.num_providers());
1187 assert_eq!(manager.num_heartbeat_workers(), 2);
1188
1189 let impl_ref = manager.worker_manager.read();
1190 assert_eq!(impl_ref.shared_worker.len(), 1);
1191 assert!(impl_ref.shared_worker.contains_key("shared_namespace"));
1192
1193 let shared_worker = impl_ref.shared_worker.get("shared_namespace").unwrap();
1194 assert_eq!(shared_worker.namespace(), "shared_namespace");
1195 }
1196
1197 #[test]
1198 fn different_namespaces_get_separate_heartbeat_managers() {
1199 let manager = ClientWorkerSet::new();
1200 let worker1 = new_mock_provider_with_heartbeat(
1201 "namespace1".to_string(),
1202 "queue1".to_string(),
1203 true,
1204 None,
1205 );
1206 let worker2 = new_mock_provider_with_heartbeat(
1207 "namespace2".to_string(),
1208 "queue1".to_string(),
1209 true,
1210 None,
1211 );
1212
1213 manager.register_worker(Arc::new(worker1), false).unwrap();
1214 manager.register_worker(Arc::new(worker2), false).unwrap();
1215
1216 assert_eq!(2, manager.num_providers());
1217 assert_eq!(manager.num_heartbeat_workers(), 2);
1218
1219 let impl_ref = manager.worker_manager.read();
1220 assert_eq!(impl_ref.num_heartbeat_workers(), 2);
1221 assert!(impl_ref.shared_worker.contains_key("namespace1"));
1222 assert!(impl_ref.shared_worker.contains_key("namespace2"));
1223 }
1224
1225 #[test]
1226 fn unregister_heartbeat_workers_cleans_up_shared_worker_when_last_removed() {
1227 let manager = ClientWorkerSet::new();
1228
1229 let worker1 = new_mock_provider_with_heartbeat(
1231 "test_namespace".to_string(),
1232 "queue1".to_string(),
1233 true,
1234 None,
1235 );
1236 let worker2 = new_mock_provider_with_heartbeat(
1237 "test_namespace".to_string(),
1238 "queue2".to_string(),
1239 true,
1240 None,
1241 );
1242 let worker_instance_key1 = worker1.worker_instance_key();
1243 let worker_instance_key2 = worker2.worker_instance_key();
1244
1245 assert_ne!(worker_instance_key1, worker_instance_key2);
1246
1247 manager.register_worker(Arc::new(worker1), false).unwrap();
1248 manager.register_worker(Arc::new(worker2), false).unwrap();
1249
1250 assert_eq!(2, manager.num_providers());
1252 assert_eq!(manager.num_heartbeat_workers(), 2);
1253
1254 let impl_ref = manager.worker_manager.read();
1255 assert_eq!(impl_ref.shared_worker.len(), 1);
1256 assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1257 assert_eq!(
1258 impl_ref
1259 .shared_worker
1260 .get("test_namespace")
1261 .unwrap()
1262 .num_workers(),
1263 2
1264 );
1265 drop(impl_ref);
1266
1267 manager
1269 .unregister_slot_provider(worker_instance_key1)
1270 .unwrap();
1271 manager.finalize_unregister(worker_instance_key1).unwrap();
1272
1273 assert_eq!(1, manager.num_providers());
1275 assert_eq!(manager.num_heartbeat_workers(), 1);
1276
1277 let impl_ref = manager.worker_manager.read();
1278 assert_eq!(impl_ref.num_heartbeat_workers(), 1); assert!(impl_ref.shared_worker.contains_key("test_namespace"));
1280 assert_eq!(
1281 impl_ref
1282 .shared_worker
1283 .get("test_namespace")
1284 .unwrap()
1285 .num_workers(),
1286 1
1287 );
1288 drop(impl_ref);
1289
1290 manager
1292 .unregister_slot_provider(worker_instance_key2)
1293 .unwrap();
1294 manager.finalize_unregister(worker_instance_key2).unwrap();
1295
1296 assert_eq!(0, manager.num_providers());
1298 assert_eq!(manager.num_heartbeat_workers(), 0);
1299
1300 let impl_ref = manager.worker_manager.read();
1301 assert_eq!(impl_ref.shared_worker.len(), 0); assert!(!impl_ref.shared_worker.contains_key("test_namespace"));
1303 }
1304
1305 #[test]
1306 fn workflow_and_activity_only_workers_coexist() {
1307 let manager = ClientWorkerSet::new();
1308 let namespace = "test_namespace".to_string();
1309 let task_queue = "test_queue".to_string();
1310
1311 let mut workflow_nexus_worker = MockClientWorker::new();
1312 workflow_nexus_worker
1313 .expect_namespace()
1314 .return_const(namespace.clone());
1315 workflow_nexus_worker
1316 .expect_task_queue()
1317 .return_const(task_queue.clone());
1318 workflow_nexus_worker
1319 .expect_deployment_options()
1320 .return_const(None);
1321 workflow_nexus_worker
1322 .expect_worker_instance_key()
1323 .return_const(Uuid::new_v4());
1324 workflow_nexus_worker
1325 .expect_heartbeat_enabled()
1326 .return_const(false);
1327 workflow_nexus_worker
1328 .expect_worker_task_types()
1329 .return_const(WorkerTaskTypes {
1330 enable_workflows: true,
1331 enable_local_activities: false,
1332 enable_remote_activities: false,
1333 enable_nexus: true,
1334 });
1335
1336 let mut activity_worker = MockClientWorker::new();
1337 activity_worker
1338 .expect_namespace()
1339 .return_const(namespace.clone());
1340 activity_worker
1341 .expect_task_queue()
1342 .return_const(task_queue.clone());
1343 activity_worker
1344 .expect_deployment_options()
1345 .return_const(None);
1346 activity_worker
1347 .expect_worker_instance_key()
1348 .return_const(Uuid::new_v4());
1349 activity_worker
1350 .expect_heartbeat_enabled()
1351 .return_const(false);
1352 activity_worker
1353 .expect_worker_task_types()
1354 .return_const(WorkerTaskTypes {
1355 enable_workflows: false,
1356 enable_local_activities: false,
1357 enable_remote_activities: true,
1358 enable_nexus: false,
1359 });
1360 activity_worker.expect_try_reserve_wft_slot().times(0); manager
1363 .register_worker(Arc::new(workflow_nexus_worker), false)
1364 .expect("workflow-nexus worker should register");
1365 manager
1366 .register_worker(Arc::new(activity_worker), false)
1367 .expect("activity-only worker should register");
1368
1369 assert_eq!(2, manager.num_providers());
1370 }
1371
1372 #[test]
1373 fn overlapping_capabilities_rejected() {
1374 let manager = ClientWorkerSet::new();
1375 let namespace = "test_namespace".to_string();
1376 let task_queue = "test_queue".to_string();
1377
1378 let mut worker1 = MockClientWorker::new();
1380 worker1.expect_namespace().return_const(namespace.clone());
1381 worker1.expect_task_queue().return_const(task_queue.clone());
1382 worker1.expect_deployment_options().return_const(None);
1383 worker1
1384 .expect_worker_instance_key()
1385 .return_const(Uuid::new_v4());
1386 worker1.expect_heartbeat_enabled().return_const(false);
1387 worker1
1388 .expect_worker_task_types()
1389 .return_const(WorkerTaskTypes {
1390 enable_workflows: true,
1391 enable_local_activities: true,
1392 enable_remote_activities: true,
1393 enable_nexus: false,
1394 });
1395
1396 let mut worker2 = MockClientWorker::new();
1398 worker2.expect_namespace().return_const(namespace.clone());
1399 worker2.expect_task_queue().return_const(task_queue.clone());
1400 worker2.expect_deployment_options().return_const(None);
1401 worker2
1402 .expect_worker_instance_key()
1403 .return_const(Uuid::new_v4());
1404 worker2.expect_heartbeat_enabled().return_const(false);
1405 worker2
1406 .expect_worker_task_types()
1407 .return_const(WorkerTaskTypes {
1408 enable_workflows: true,
1409 enable_local_activities: true,
1410 enable_remote_activities: true,
1411 enable_nexus: false,
1412 });
1413
1414 manager
1415 .register_worker(Arc::new(worker1), false)
1416 .expect("first worker should register");
1417
1418 let result = manager.register_worker(Arc::new(worker2), false);
1419 assert!(result.is_err());
1420 assert!(
1421 result
1422 .unwrap_err()
1423 .to_string()
1424 .contains("overlapping worker task types")
1425 );
1426
1427 let mut worker3 = MockClientWorker::new();
1429 worker3.expect_namespace().return_const(namespace.clone());
1430 worker3.expect_task_queue().return_const(task_queue.clone());
1431 worker3.expect_deployment_options().return_const(None);
1432 worker3
1433 .expect_worker_instance_key()
1434 .return_const(Uuid::new_v4());
1435 worker3.expect_heartbeat_enabled().return_const(false);
1436 worker3
1437 .expect_worker_task_types()
1438 .return_const(WorkerTaskTypes {
1439 enable_workflows: false,
1440 enable_local_activities: false,
1441 enable_remote_activities: true,
1442 enable_nexus: false,
1443 });
1444
1445 let result = manager.register_worker(Arc::new(worker3), false);
1446 assert!(result.is_err());
1447 assert!(
1448 result
1449 .unwrap_err()
1450 .to_string()
1451 .contains("overlapping worker task types")
1452 );
1453 }
1454
1455 #[test]
1456 fn wft_slot_reservation_ignores_non_workflow_workers() {
1457 let mut manager_impl = ClientWorkerSetImpl::new();
1458 let namespace = "test_namespace".to_string();
1459 let task_queue = "test_queue".to_string();
1460
1461 let mut activity_worker = MockClientWorker::new();
1462 activity_worker
1463 .expect_namespace()
1464 .return_const(namespace.clone());
1465 activity_worker
1466 .expect_task_queue()
1467 .return_const(task_queue.clone());
1468 activity_worker
1469 .expect_deployment_options()
1470 .return_const(None);
1471 activity_worker
1472 .expect_worker_instance_key()
1473 .return_const(Uuid::new_v4());
1474 activity_worker
1475 .expect_heartbeat_enabled()
1476 .return_const(false);
1477 activity_worker
1478 .expect_worker_task_types()
1479 .return_const(WorkerTaskTypes {
1480 enable_workflows: false,
1481 enable_local_activities: false,
1482 enable_remote_activities: true,
1483 enable_nexus: false,
1484 });
1485
1486 let mut nexus_worker = MockClientWorker::new();
1487 nexus_worker
1488 .expect_namespace()
1489 .return_const(namespace.clone());
1490 nexus_worker
1491 .expect_task_queue()
1492 .return_const(task_queue.clone());
1493 nexus_worker.expect_deployment_options().return_const(None);
1494 nexus_worker
1495 .expect_worker_instance_key()
1496 .return_const(Uuid::new_v4());
1497 nexus_worker.expect_heartbeat_enabled().return_const(false);
1498 nexus_worker
1499 .expect_worker_task_types()
1500 .return_const(WorkerTaskTypes {
1501 enable_workflows: false,
1502 enable_local_activities: false,
1503 enable_remote_activities: false,
1504 enable_nexus: true,
1505 });
1506
1507 manager_impl
1508 .register(Arc::new(activity_worker), false)
1509 .expect("activity worker should register");
1510 manager_impl
1511 .register(Arc::new(nexus_worker), false)
1512 .expect("nexus worker should register");
1513
1514 let reservation = manager_impl.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1515 assert!(
1516 reservation.is_none(),
1517 "should not find workflow workers when only activity/nexus workers registered"
1518 );
1519
1520 let mut workflow_worker = MockClientWorker::new();
1522 workflow_worker
1523 .expect_namespace()
1524 .return_const(namespace.clone());
1525 workflow_worker
1526 .expect_task_queue()
1527 .return_const(task_queue.clone());
1528 workflow_worker
1529 .expect_deployment_options()
1530 .return_const(None);
1531 workflow_worker
1532 .expect_worker_instance_key()
1533 .return_const(Uuid::new_v4());
1534 workflow_worker
1535 .expect_heartbeat_enabled()
1536 .return_const(false);
1537 workflow_worker
1538 .expect_worker_task_types()
1539 .return_const(WorkerTaskTypes {
1540 enable_workflows: true,
1541 enable_local_activities: true,
1542 enable_remote_activities: false,
1543 enable_nexus: false,
1544 });
1545 workflow_worker
1546 .expect_try_reserve_wft_slot()
1547 .times(1)
1548 .returning(|| Some(new_mock_slot(false)));
1549
1550 manager_impl
1551 .register(Arc::new(workflow_worker), false)
1552 .expect("workflow worker should register");
1553
1554 let reservation = manager_impl.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1555 assert!(
1556 reservation.is_some(),
1557 "should find workflow worker after it's registered"
1558 );
1559 }
1560
1561 #[test]
1562 fn worker_invalid_type_config_rejected() {
1563 let manager = ClientWorkerSet::new();
1564
1565 let mut worker = MockClientWorker::new();
1567 worker
1568 .expect_namespace()
1569 .return_const("test_namespace".to_string());
1570 worker
1571 .expect_task_queue()
1572 .return_const("test_queue".to_string());
1573 worker.expect_deployment_options().return_const(None);
1574 worker
1575 .expect_worker_instance_key()
1576 .return_const(Uuid::new_v4());
1577 worker.expect_heartbeat_enabled().return_const(false);
1578 worker
1579 .expect_worker_task_types()
1580 .return_const(WorkerTaskTypes {
1581 enable_workflows: false,
1582 enable_local_activities: false,
1583 enable_remote_activities: false,
1584 enable_nexus: false,
1585 });
1586
1587 let result = manager.register_worker(Arc::new(worker), false);
1588 assert!(result.is_err());
1589 assert!(
1590 result
1591 .unwrap_err()
1592 .to_string()
1593 .contains("must have at least one capability enabled")
1594 );
1595
1596 let mut worker = MockClientWorker::new();
1598 worker
1599 .expect_namespace()
1600 .return_const("test_namespace".to_string());
1601 worker
1602 .expect_task_queue()
1603 .return_const("test_queue".to_string());
1604 worker.expect_deployment_options().return_const(None);
1605 worker
1606 .expect_worker_instance_key()
1607 .return_const(Uuid::new_v4());
1608 worker.expect_heartbeat_enabled().return_const(false);
1609 worker
1610 .expect_worker_task_types()
1611 .return_const(WorkerTaskTypes {
1612 enable_workflows: false,
1613 enable_local_activities: true,
1614 enable_remote_activities: true,
1615 enable_nexus: false,
1616 });
1617
1618 let result = manager.register_worker(Arc::new(worker), false);
1619 assert!(result.is_err());
1620 assert_eq!(
1621 result.unwrap_err().to_string(),
1622 "Local activities cannot be enabled without workflows".to_string()
1623 );
1624 }
1625
1626 #[test]
1627 fn unregister_with_multiple_workers() {
1628 let manager = ClientWorkerSet::new();
1629 let namespace = "test_namespace".to_string();
1630 let task_queue = "test_queue".to_string();
1631
1632 let mut workflow_worker = MockClientWorker::new();
1634 workflow_worker
1635 .expect_namespace()
1636 .return_const(namespace.clone());
1637 workflow_worker
1638 .expect_task_queue()
1639 .return_const(task_queue.clone());
1640 workflow_worker
1641 .expect_deployment_options()
1642 .return_const(None);
1643 let wf_worker_key = Uuid::new_v4();
1644 workflow_worker
1645 .expect_worker_instance_key()
1646 .return_const(wf_worker_key);
1647 workflow_worker
1648 .expect_heartbeat_enabled()
1649 .return_const(false);
1650 workflow_worker
1651 .expect_worker_task_types()
1652 .return_const(WorkerTaskTypes {
1653 enable_workflows: true,
1654 enable_local_activities: true,
1655 enable_remote_activities: false,
1656 enable_nexus: false,
1657 });
1658 workflow_worker
1659 .expect_try_reserve_wft_slot()
1660 .returning(|| Some(new_mock_slot(false)));
1661
1662 let mut activity_worker = MockClientWorker::new();
1664 activity_worker
1665 .expect_namespace()
1666 .return_const(namespace.clone());
1667 activity_worker
1668 .expect_task_queue()
1669 .return_const(task_queue.clone());
1670 activity_worker
1671 .expect_deployment_options()
1672 .return_const(None);
1673 let act_worker_key = Uuid::new_v4();
1674 activity_worker
1675 .expect_worker_instance_key()
1676 .return_const(act_worker_key);
1677 activity_worker
1678 .expect_heartbeat_enabled()
1679 .return_const(false);
1680 activity_worker
1681 .expect_worker_task_types()
1682 .return_const(WorkerTaskTypes {
1683 enable_workflows: false,
1684 enable_local_activities: false,
1685 enable_remote_activities: true,
1686 enable_nexus: false,
1687 });
1688
1689 manager
1690 .register_worker(Arc::new(workflow_worker), false)
1691 .expect("workflow worker should register");
1692 manager
1693 .register_worker(Arc::new(activity_worker), false)
1694 .expect("activity worker should register");
1695
1696 assert_eq!(2, manager.num_providers());
1697
1698 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1699 assert!(
1700 reservation.is_some(),
1701 "should be able to reserve slot from workflow worker"
1702 );
1703
1704 manager
1705 .unregister_slot_provider(wf_worker_key)
1706 .expect("should unregister slot provider for workflow worker");
1707 manager
1708 .finalize_unregister(wf_worker_key)
1709 .expect("should finalize unregister for workflow worker");
1710
1711 assert_eq!(1, manager.num_providers());
1713
1714 let reservation = manager.try_reserve_wft_slot(namespace.clone(), task_queue.clone());
1715 assert!(
1716 reservation.is_none(),
1717 "should not find workflow worker after unregistration"
1718 );
1719
1720 manager
1721 .unregister_slot_provider(act_worker_key)
1722 .expect("should unregister slot provider for activity worker");
1723 manager
1724 .finalize_unregister(act_worker_key)
1725 .expect("should finalize unregister for activity worker");
1726
1727 assert_eq!(0, manager.num_providers());
1728 }
1729
1730 #[test]
1731 fn worker_unregister_order() {
1732 let manager = ClientWorkerSet::new();
1733 let worker = new_mock_provider_with_heartbeat(
1734 "namespace1".to_string(),
1735 "queue1".to_string(),
1736 true,
1737 None,
1738 );
1739 let worker_instance_key = worker.worker_instance_key();
1740 manager.register_worker(Arc::new(worker), false).unwrap();
1741
1742 let res = manager.finalize_unregister(worker_instance_key);
1743 assert!(res.is_err());
1744 let err_string = res.err().map(|e| e.to_string()).unwrap();
1745 assert!(err_string.contains("Worker still in slot_providers during finalize"));
1746
1747 manager
1750 .unregister_slot_provider(worker_instance_key)
1751 .unwrap();
1752 manager.finalize_unregister(worker_instance_key).unwrap();
1753 }
1754}