Skip to main content

temporalio_client/
worker.rs

1//! Contains types and logic for interactions between clients and Core/SDK workers.
2//!
3//! This module is public for use by `temporalio-sdk-core` and is not intended to be used directly
4//! by SDK users.
5
6use 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/// This trait represents a slot reserved for processing a WFT by a worker.
32#[cfg_attr(test, mockall::automock)]
33pub trait Slot {
34    /// Consumes this slot by dispatching a WFT to its worker. This can only be called once.
35    fn schedule_wft(
36        self: Box<Self>,
37        task: PollWorkflowTaskQueueResponse,
38    ) -> Result<(), anyhow::Error>;
39}
40
41/// Result of reserving a workflow task slot, including deployment options if applicable.
42pub(crate) struct SlotReservation {
43    /// The reserved slot for processing the workflow task
44    pub slot: Box<dyn Slot + Send>,
45    /// Worker deployment options, if the worker is using deployment-based versioning
46    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/// Information about a registered worker in the slot provider registry
65#[derive(Debug, Clone)]
66struct RegisteredWorkerInfo {
67    /// Unique identifier for this worker instance
68    worker_id: Uuid,
69    /// Optional deployment build ID for versioning
70    build_id: Option<String>,
71    /// Task types this worker can handle
72    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
85/// This is an inner class for [ClientWorkerSet] needed to hide the mutex.
86struct ClientWorkerSetImpl {
87    /// Maps slot keys to registered worker information
88    slot_providers: HashMap<SlotKey, Vec<RegisteredWorkerInfo>>,
89    /// Maps worker_instance_key to registered workers
90    all_workers: HashMap<Uuid, Arc<dyn ClientWorker + Send + Sync>>,
91    /// Maps namespace to shared worker for worker heartbeating
92    shared_worker: HashMap<String, Box<dyn SharedNamespaceWorkerTrait + Send + Sync>>,
93    // Avoid retaining namespace limits and capabilities after the last worker using them is gone.
94    namespace_descriptions: HashMap<String, Weak<NamespaceDescriptionSource>>,
95}
96
97impl ClientWorkerSetImpl {
98    /// Factory method.
99    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        // For tests we return workers in the order they're registered, so we can test
158        // the retry mechanism deterministically
159        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    /// Slot provider should be unregistered at the beginning of worker shutdown, in order to disable
258    /// eager workflow start.
259    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/// A connection-scoped source for a namespace description shared by all workers in that namespace.
322#[derive(Debug)]
323#[doc(hidden)]
324pub struct NamespaceDescriptionSource {
325    description: OnceCell<DescribeNamespaceResponse>,
326}
327
328impl NamespaceDescriptionSource {
329    /// Construct a source whose description has not yet been resolved.
330    pub fn unresolved() -> Self {
331        Self {
332            description: OnceCell::new(),
333        }
334    }
335
336    /// Construct a source whose description has already been resolved.
337    pub fn resolved(description: DescribeNamespaceResponse) -> Self {
338        Self {
339            description: OnceCell::new_with(Some(description)),
340        }
341    }
342
343    /// Resolve the namespace description once, allowing concurrent callers to await the same RPC.
344    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    /// Return the resolved namespace description, if resolution has completed successfully.
365    pub fn get(&self) -> Option<&DescribeNamespaceResponse> {
366        self.description.get()
367    }
368}
369
370/// This trait represents a shared namespace worker that sends worker heartbeats and
371/// receives worker commands.
372pub trait SharedNamespaceWorkerTrait {
373    /// Namespace that the shared namespace worker is connected to.
374    fn namespace(&self) -> String;
375
376    /// Registers worker callbacks.
377    fn register_callback(&self, worker_instance_key: Uuid, callbacks: WorkerCallbacks);
378
379    /// Unregisters worker callbacks. Returns the callbacks removed, as well as a bool that
380    /// indicates if there are no remaining callbacks in the SharedNamespaceWorker, indicating
381    /// the shared worker itself can be shut down.
382    fn unregister_callback(&self, worker_instance_key: Uuid) -> (Option<WorkerCallbacks>, bool);
383
384    /// Returns the number of workers registered to this shared worker.
385    fn num_workers(&self) -> usize;
386
387    /// Returns whether this shared worker is polling the worker control task queue.
388    fn worker_control_task_queue_enabled(&self) -> bool {
389        false
390    }
391}
392
393/// Enables local workers to make themselves visible to a shared client instance.
394///
395/// For slot managing, there can only be one worker registered per
396/// namespace+queue_name+connection, others will return an error.
397/// It also provides a convenient method to find compatible slots within the collection.
398pub 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    /// Factory method.
411    pub fn new() -> Self {
412        Self {
413            worker_grouping_key: Uuid::new_v4(),
414            worker_manager: RwLock::new(ClientWorkerSetImpl::new()),
415        }
416    }
417
418    /// Return the shared namespace description source for this connection and namespace.
419    #[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    /// Try to reserve a compatible processing slot in any of the registered workers.
427    /// Returns the slot and the worker's deployment options (if using deployment-based versioning).
428    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    /// Register a local worker that can provide WFT processing slots and potentially worker heartbeating.
439    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    /// Disables Eager Workflow Start for this worker. This must be called before
450    /// `finalize_unregister`, otherwise `finalize_unregister` will return an err.
451    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    /// Finalizes unregistering of worker from client. This must be called at the end of worker
458    /// shutdown in order to finalize shutdown for worker heartbeat properly. Must call after
459    /// `unregister_slot_provider`, otherwise an err will be returned.
460    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    /// Returns the worker grouping key, which is unique for each worker.
470    pub fn worker_grouping_key(&self) -> Uuid {
471        self.worker_grouping_key
472    }
473
474    /// Returns whether the shared worker for `namespace` is polling the worker control task queue.
475    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    /// Returns (num_providers, num_buckets), where a bucket key is namespace+task_queue.
483    /// There is only one provider per bucket so `num_providers` should be equal to `num_buckets`.
484    pub fn num_providers(&self) -> usize {
485        self.worker_manager.read().num_providers()
486    }
487
488    #[cfg(test)]
489    /// Returns the total number of heartbeat workers registered across all namespaces.
490    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
503/// Contains a worker heartbeat callback, wrapped for mocking. Returns `None` until the worker has
504/// started running.
505pub type HeartbeatCallback = Arc<dyn Fn() -> Option<WorkerHeartbeat> + Send + Sync>;
506
507/// Callback invoked after a worker heartbeat has been accepted by the server.
508pub type HeartbeatSuccessCallback = Arc<dyn Fn() + Send + Sync>;
509
510/// Callback to cancel an activity by task token. Returns true if the activity was found.
511pub type CancelActivityCallback = Arc<dyn Fn(TaskToken) -> bool + Send + Sync>;
512
513/// Bundles all per-worker callbacks registered with the SharedNamespaceWorker.
514pub struct WorkerCallbacks {
515    /// Callback to collect heartbeat data from the worker.
516    pub heartbeat: HeartbeatCallback,
517    /// Callback acknowledging successful delivery of the collected heartbeat.
518    pub heartbeat_success: Option<HeartbeatSuccessCallback>,
519    /// Callback to cancel an activity by task token.
520    pub cancel_activity: Option<CancelActivityCallback>,
521}
522
523/// Represents a complete worker that can handle both slot management
524/// and worker heartbeat functionality.
525#[cfg_attr(test, mockall::automock)]
526pub trait ClientWorker: Send + Sync {
527    /// The namespace this worker operates in
528    fn namespace(&self) -> &str;
529
530    /// The task queue this worker listens to
531    fn task_queue(&self) -> &str;
532
533    /// Try to reserve a slot for workflow task processing.
534    ///
535    /// This method should return `Some(slot)` if a workflow task slot is available,
536    /// or `None` if all slots are currently in use. The returned slot will be used
537    /// to process exactly one workflow task.
538    fn try_reserve_wft_slot(&self) -> Option<Box<dyn Slot + Send>>;
539
540    /// Get the worker deployment options for this worker, if using deployment-based versioning.
541    fn deployment_options(&self) -> Option<WorkerDeploymentOptions>;
542
543    /// Unique identifier for this worker instance.
544    /// This must be stable across the worker's lifetime and unique per instance.
545    fn worker_instance_key(&self) -> Uuid;
546
547    /// Indicates if worker heartbeating is enabled for this client worker.
548    fn heartbeat_enabled(&self) -> bool;
549
550    /// Returns the heartbeat callback that can be used to get WorkerHeartbeat data.
551    fn heartbeat_callback(&self) -> Option<HeartbeatCallback>;
552
553    /// Returns a callback notified after the heartbeat is accepted by the server.
554    fn heartbeat_success_callback(&self) -> Option<HeartbeatSuccessCallback> {
555        None
556    }
557
558    /// Returns a callback that can cancel an activity by task token.
559    fn cancel_activity_callback(&self) -> Option<CancelActivityCallback>;
560
561    /// Creates a new worker that implements the [SharedNamespaceWorkerTrait]
562    fn new_shared_namespace_worker(
563        &self,
564    ) -> Result<Box<dyn SharedNamespaceWorkerTrait + Send + Sync>, anyhow::Error>;
565
566    /// Returns the task types this worker can handle
567    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        // On a separate task queue
811        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                // Should get error for overlapping worker task types
867                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            // expect error since worker is already unregistered
883            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        // Same namespace+task_queue but different worker instance
1054        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        // second worker register should fail due to overlapping worker task types
1064        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        // Same namespace but different task queue
1176        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        // Create two workers with same namespace but different task queues
1230        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        // Verify initial state: 2 slot providers, 2 heartbeat workers, 1 shared worker
1251        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        // Unregister first worker
1268        manager
1269            .unregister_slot_provider(worker_instance_key1)
1270            .unwrap();
1271        manager.finalize_unregister(worker_instance_key1).unwrap();
1272
1273        // After unregistering first worker: 1 slot provider, 1 heartbeat worker, shared worker still exists
1274        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); // SharedNamespaceWorker still exists
1279        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        // Unregister second worker
1291        manager
1292            .unregister_slot_provider(worker_instance_key2)
1293            .unwrap();
1294        manager.finalize_unregister(worker_instance_key2).unwrap();
1295
1296        // After unregistering last worker: 0 slot providers, 0 heartbeat workers, shared worker is removed
1297        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); // SharedNamespaceWorker is cleaned up
1302        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); // Should not be called for activity-only worker
1361
1362        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        // workflow+activity worker
1379        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        // workflow+activity worker
1397        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        // activity-only worker
1428        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        // Now register a workflow worker
1521        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        // no types enabled
1566        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        // local activities enabled without workflows
1597        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        // workflow-only worker
1633        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        // activity-only worker
1663        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        // Activity worker should still be registered
1712        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        // previous incorrect call to finalize_unregister should not cause any state leaks when
1748        // properly removed later
1749        manager
1750            .unregister_slot_provider(worker_instance_key)
1751            .unwrap();
1752        manager.finalize_unregister(worker_instance_key).unwrap();
1753    }
1754}