Skip to main content

temporalio_client/
worker.rs

1//! Contains types and logic for interactions between clients and Core/SDK workers
2
3use 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/// This trait represents a slot reserved for processing a WFT by a worker.
29#[cfg_attr(test, mockall::automock)]
30pub trait Slot {
31    /// Consumes this slot by dispatching a WFT to its worker. This can only be called once.
32    fn schedule_wft(
33        self: Box<Self>,
34        task: PollWorkflowTaskQueueResponse,
35    ) -> Result<(), anyhow::Error>;
36}
37
38/// Result of reserving a workflow task slot, including deployment options if applicable.
39pub(crate) struct SlotReservation {
40    /// The reserved slot for processing the workflow task
41    pub slot: Box<dyn Slot + Send>,
42    /// Worker deployment options, if the worker is using deployment-based versioning
43    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/// Information about a registered worker in the slot provider registry
62#[derive(Debug, Clone)]
63struct RegisteredWorkerInfo {
64    /// Unique identifier for this worker instance
65    worker_id: Uuid,
66    /// Optional deployment build ID for versioning
67    build_id: Option<String>,
68    /// Task types this worker can handle
69    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
82/// This is an inner class for [ClientWorkerSet] needed to hide the mutex.
83struct ClientWorkerSetImpl {
84    /// Maps slot keys to registered worker information
85    slot_providers: HashMap<SlotKey, Vec<RegisteredWorkerInfo>>,
86    /// Maps worker_instance_key to registered workers
87    all_workers: HashMap<Uuid, Arc<dyn ClientWorker + Send + Sync>>,
88    /// Maps namespace to shared worker for worker heartbeating
89    shared_worker: HashMap<String, Box<dyn SharedNamespaceWorkerTrait + Send + Sync>>,
90    // Avoid retaining namespace limits and capabilities after the last worker using them is gone.
91    namespace_descriptions: HashMap<String, Weak<NamespaceDescriptionSource>>,
92}
93
94impl ClientWorkerSetImpl {
95    /// Factory method.
96    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        // For tests we return workers in the order they're registered, so we can test
155        // the retry mechanism deterministically
156        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    /// Slot provider should be unregistered at the beginning of worker shutdown, in order to disable
254    /// eager workflow start.
255    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/// A connection-scoped source for a namespace description shared by all workers in that namespace.
318#[derive(Debug)]
319#[doc(hidden)]
320pub struct NamespaceDescriptionSource {
321    description: OnceCell<DescribeNamespaceResponse>,
322}
323
324impl NamespaceDescriptionSource {
325    /// Construct a source whose description has not yet been resolved.
326    pub fn unresolved() -> Self {
327        Self {
328            description: OnceCell::new(),
329        }
330    }
331
332    /// Construct a source whose description has already been resolved.
333    pub fn resolved(description: DescribeNamespaceResponse) -> Self {
334        Self {
335            description: OnceCell::new_with(Some(description)),
336        }
337    }
338
339    /// Resolve the namespace description once, allowing concurrent callers to await the same RPC.
340    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    /// Return the resolved namespace description, if resolution has completed successfully.
361    pub fn get(&self) -> Option<&DescribeNamespaceResponse> {
362        self.description.get()
363    }
364}
365
366/// This trait represents a shared namespace worker that sends worker heartbeats and
367/// receives worker commands.
368pub trait SharedNamespaceWorkerTrait {
369    /// Namespace that the shared namespace worker is connected to.
370    fn namespace(&self) -> String;
371
372    /// Registers worker callbacks.
373    fn register_callback(&self, worker_instance_key: Uuid, callbacks: WorkerCallbacks);
374
375    /// Unregisters worker callbacks. Returns the callbacks removed, as well as a bool that
376    /// indicates if there are no remaining callbacks in the SharedNamespaceWorker, indicating
377    /// the shared worker itself can be shut down.
378    fn unregister_callback(&self, worker_instance_key: Uuid) -> (Option<WorkerCallbacks>, bool);
379
380    /// Returns the number of workers registered to this shared worker.
381    fn num_workers(&self) -> usize;
382
383    /// Returns whether this shared worker is polling the worker control task queue.
384    fn worker_control_task_queue_enabled(&self) -> bool {
385        false
386    }
387}
388
389/// Enables local workers to make themselves visible to a shared client instance.
390///
391/// For slot managing, there can only be one worker registered per
392/// namespace+queue_name+connection, others will return an error.
393/// It also provides a convenient method to find compatible slots within the collection.
394pub 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    /// Factory method.
407    pub fn new() -> Self {
408        Self {
409            worker_grouping_key: Uuid::new_v4(),
410            worker_manager: RwLock::new(ClientWorkerSetImpl::new()),
411        }
412    }
413
414    /// Return the shared namespace description source for this connection and namespace.
415    #[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    /// Try to reserve a compatible processing slot in any of the registered workers.
423    /// Returns the slot and the worker's deployment options (if using deployment-based versioning).
424    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    /// Register a local worker that can provide WFT processing slots and potentially worker heartbeating.
435    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    /// Disables Eager Workflow Start for this worker. This must be called before
446    /// `finalize_unregister`, otherwise `finalize_unregister` will return an err.
447    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    /// Finalizes unregistering of worker from client. This must be called at the end of worker
454    /// shutdown in order to finalize shutdown for worker heartbeat properly. Must call after
455    /// `unregister_slot_provider`, otherwise an err will be returned.
456    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    /// Returns the worker grouping key, which is unique for each worker.
466    pub fn worker_grouping_key(&self) -> Uuid {
467        self.worker_grouping_key
468    }
469
470    /// Returns whether the shared worker for `namespace` is polling the worker control task queue.
471    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    /// Returns (num_providers, num_buckets), where a bucket key is namespace+task_queue.
479    /// There is only one provider per bucket so `num_providers` should be equal to `num_buckets`.
480    pub fn num_providers(&self) -> usize {
481        self.worker_manager.read().num_providers()
482    }
483
484    #[cfg(test)]
485    /// Returns the total number of heartbeat workers registered across all namespaces.
486    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
499/// Contains a worker heartbeat callback, wrapped for mocking
500pub type HeartbeatCallback = Arc<dyn Fn() -> WorkerHeartbeat + Send + Sync>;
501
502/// Callback to cancel an activity by task token. Returns true if the activity was found.
503pub type CancelActivityCallback = Arc<dyn Fn(TaskToken) -> bool + Send + Sync>;
504
505/// Bundles all per-worker callbacks registered with the SharedNamespaceWorker.
506pub struct WorkerCallbacks {
507    /// Callback to collect heartbeat data from the worker.
508    pub heartbeat: HeartbeatCallback,
509    /// Callback to cancel an activity by task token.
510    pub cancel_activity: Option<CancelActivityCallback>,
511}
512
513/// Represents a complete worker that can handle both slot management
514/// and worker heartbeat functionality.
515#[cfg_attr(test, mockall::automock)]
516pub trait ClientWorker: Send + Sync {
517    /// The namespace this worker operates in
518    fn namespace(&self) -> &str;
519
520    /// The task queue this worker listens to
521    fn task_queue(&self) -> &str;
522
523    /// Try to reserve a slot for workflow task processing.
524    ///
525    /// This method should return `Some(slot)` if a workflow task slot is available,
526    /// or `None` if all slots are currently in use. The returned slot will be used
527    /// to process exactly one workflow task.
528    fn try_reserve_wft_slot(&self) -> Option<Box<dyn Slot + Send>>;
529
530    /// Get the worker deployment options for this worker, if using deployment-based versioning.
531    fn deployment_options(&self) -> Option<WorkerDeploymentOptions>;
532
533    /// Unique identifier for this worker instance.
534    /// This must be stable across the worker's lifetime and unique per instance.
535    fn worker_instance_key(&self) -> Uuid;
536
537    /// Indicates if worker heartbeating is enabled for this client worker.
538    fn heartbeat_enabled(&self) -> bool;
539
540    /// Returns the heartbeat callback that can be used to get WorkerHeartbeat data.
541    fn heartbeat_callback(&self) -> Option<HeartbeatCallback>;
542
543    /// Returns a callback that can cancel an activity by task token.
544    fn cancel_activity_callback(&self) -> Option<CancelActivityCallback>;
545
546    /// Creates a new worker that implements the [SharedNamespaceWorkerTrait]
547    fn new_shared_namespace_worker(
548        &self,
549    ) -> Result<Box<dyn SharedNamespaceWorkerTrait + Send + Sync>, anyhow::Error>;
550
551    /// Returns the task types this worker can handle
552    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        // On a separate task queue
801        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                // Should get error for overlapping worker task types
857                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            // expect error since worker is already unregistered
873            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        // Same namespace+task_queue but different worker instance
1041        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        // second worker register should fail due to overlapping worker task types
1051        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        // Same namespace but different task queue
1163        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        // Create two workers with same namespace but different task queues
1217        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        // Verify initial state: 2 slot providers, 2 heartbeat workers, 1 shared worker
1238        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        // Unregister first worker
1255        manager
1256            .unregister_slot_provider(worker_instance_key1)
1257            .unwrap();
1258        manager.finalize_unregister(worker_instance_key1).unwrap();
1259
1260        // After unregistering first worker: 1 slot provider, 1 heartbeat worker, shared worker still exists
1261        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); // SharedNamespaceWorker still exists
1266        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        // Unregister second worker
1278        manager
1279            .unregister_slot_provider(worker_instance_key2)
1280            .unwrap();
1281        manager.finalize_unregister(worker_instance_key2).unwrap();
1282
1283        // After unregistering last worker: 0 slot providers, 0 heartbeat workers, shared worker is removed
1284        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); // SharedNamespaceWorker is cleaned up
1289        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); // Should not be called for activity-only worker
1348
1349        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        // workflow+activity worker
1366        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        // workflow+activity worker
1384        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        // activity-only worker
1415        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        // Now register a workflow worker
1508        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        // no types enabled
1553        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        // local activities enabled without workflows
1584        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        // workflow-only worker
1620        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        // activity-only worker
1650        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        // Activity worker should still be registered
1699        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        // previous incorrect call to finalize_unregister should not cause any state leaks when
1735        // properly removed later
1736        manager
1737            .unregister_slot_provider(worker_instance_key)
1738            .unwrap();
1739        manager.finalize_unregister(worker_instance_key).unwrap();
1740    }
1741}