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