Skip to main content

risingwave_batch/worker_manager/
worker_node_manager.rs

1// Copyright 2024 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{HashMap, HashSet};
16use std::sync::{Arc, RwLock};
17use std::time::Duration;
18
19use rand::seq::IndexedRandom;
20use risingwave_common::bail;
21use risingwave_common::hash::{WorkerSlotId, WorkerSlotMapping};
22use risingwave_common::id::{FragmentId, WorkerId};
23use risingwave_common::util::version::current_rw_version;
24use risingwave_common::vnode_mapping::vnode_placement::place_vnode;
25use risingwave_pb::common::{WorkerNode, WorkerType};
26
27use crate::error::{BatchError, Result};
28
29/// `WorkerNodeManager` manages live worker nodes and table vnode mapping information.
30pub struct WorkerNodeManager {
31    inner: RwLock<WorkerNodeManagerInner>,
32    /// Temporarily make worker invisible from serving cluster.
33    worker_node_mask: Arc<RwLock<HashSet<WorkerId>>>,
34}
35
36struct WorkerNodeManagerInner {
37    worker_nodes: HashMap<WorkerId, WorkerNode>,
38    /// fragment vnode mapping info for streaming
39    streaming_fragment_vnode_mapping: Option<HashMap<FragmentId, WorkerSlotMapping>>,
40    /// fragment vnode mapping info for serving
41    serving_fragment_vnode_mapping: HashMap<FragmentId, WorkerSlotMapping>,
42}
43
44pub type WorkerNodeManagerRef = Arc<WorkerNodeManager>;
45
46impl Default for WorkerNodeManager {
47    fn default() -> Self {
48        Self::new()
49    }
50}
51
52impl WorkerNodeManager {
53    pub fn new() -> Self {
54        Self {
55            inner: RwLock::new(WorkerNodeManagerInner {
56                worker_nodes: Default::default(),
57                streaming_fragment_vnode_mapping: None,
58                serving_fragment_vnode_mapping: Default::default(),
59            }),
60            worker_node_mask: Arc::new(Default::default()),
61        }
62    }
63
64    /// Used in tests.
65    pub fn mock(worker_nodes: Vec<WorkerNode>) -> Self {
66        let worker_nodes = worker_nodes.into_iter().map(|w| (w.id, w)).collect();
67        let inner = RwLock::new(WorkerNodeManagerInner {
68            worker_nodes,
69            streaming_fragment_vnode_mapping: None,
70            serving_fragment_vnode_mapping: HashMap::new(),
71        });
72        Self {
73            inner,
74            worker_node_mask: Arc::new(Default::default()),
75        }
76    }
77
78    pub fn list_compute_nodes(&self) -> Vec<WorkerNode> {
79        self.inner
80            .read()
81            .unwrap()
82            .worker_nodes
83            .values()
84            .filter(|w| w.r#type() == WorkerType::ComputeNode)
85            .cloned()
86            .collect()
87    }
88
89    pub fn list_frontend_nodes(&self) -> Vec<WorkerNode> {
90        self.inner
91            .read()
92            .unwrap()
93            .worker_nodes
94            .values()
95            .filter(|w| w.r#type() == WorkerType::Frontend)
96            .cloned()
97            .collect()
98    }
99
100    fn list_serving_worker_nodes(&self) -> Vec<WorkerNode> {
101        self.list_compute_nodes()
102            .into_iter()
103            .filter(|w| w.property.as_ref().is_some_and(|p| p.is_serving))
104            .collect()
105    }
106
107    fn list_streaming_worker_nodes(&self) -> Vec<WorkerNode> {
108        self.list_compute_nodes()
109            .into_iter()
110            .filter(|w| w.property.as_ref().is_some_and(|p| p.is_streaming))
111            .collect()
112    }
113
114    pub fn add_worker_node(&self, node: WorkerNode) {
115        let mut write_guard = self.inner.write().unwrap();
116        write_guard.worker_nodes.insert(node.id, node);
117    }
118
119    pub fn remove_worker_node(&self, node: WorkerNode) {
120        let mut write_guard = self.inner.write().unwrap();
121        write_guard.worker_nodes.remove(&node.id);
122    }
123
124    pub fn refresh(
125        &self,
126        nodes: Vec<WorkerNode>,
127        streaming_mapping: HashMap<FragmentId, WorkerSlotMapping>,
128        serving_mapping: HashMap<FragmentId, WorkerSlotMapping>,
129    ) {
130        let mut write_guard = self.inner.write().unwrap();
131        tracing::debug!("Refresh worker nodes {:?}.", nodes);
132        tracing::debug!(
133            "Refresh streaming vnode mapping for fragments {:?}.",
134            streaming_mapping.keys()
135        );
136        tracing::debug!(
137            "Refresh serving vnode mapping for fragments {:?}.",
138            serving_mapping.keys()
139        );
140        write_guard.worker_nodes = nodes.into_iter().map(|w| (w.id, w)).collect();
141        write_guard.streaming_fragment_vnode_mapping = Some(streaming_mapping);
142        write_guard.serving_fragment_vnode_mapping = serving_mapping;
143    }
144
145    /// If worker slot ids is empty, the scheduler may fail to schedule any task and stuck at
146    /// schedule next stage. If we do not return error in this case, needs more complex control
147    /// logic above. Report in this function makes the schedule root fail reason more clear.
148    pub fn get_workers_by_worker_slot_ids(
149        &self,
150        worker_slot_ids: &[WorkerSlotId],
151    ) -> Result<Vec<WorkerNode>> {
152        if worker_slot_ids.is_empty() {
153            return Err(BatchError::EmptyWorkerNodes);
154        }
155        let guard = self.inner.read().unwrap();
156        let mut workers = Vec::with_capacity(worker_slot_ids.len());
157        for worker_slot_id in worker_slot_ids {
158            match guard.worker_nodes.get(&worker_slot_id.worker_id()) {
159                Some(worker) => workers.push((*worker).clone()),
160                None => bail!(
161                    "No worker node found for worker slot id: {}",
162                    worker_slot_id
163                ),
164            }
165        }
166
167        Ok(workers)
168    }
169
170    pub fn get_streaming_fragment_mapping(
171        &self,
172        fragment_id: &FragmentId,
173    ) -> Result<WorkerSlotMapping> {
174        let guard = self.inner.read().unwrap();
175
176        let Some(streaming_mapping) = guard.streaming_fragment_vnode_mapping.as_ref() else {
177            return Err(BatchError::StreamingVnodeMappingNotInitialized);
178        };
179
180        streaming_mapping
181            .get(fragment_id)
182            .cloned()
183            .ok_or_else(|| BatchError::StreamingVnodeMappingNotFound(*fragment_id))
184    }
185
186    pub fn insert_streaming_fragment_mapping(
187        &self,
188        fragment_id: FragmentId,
189        vnode_mapping: WorkerSlotMapping,
190    ) {
191        let mut guard = self.inner.write().unwrap();
192        let mapping = guard
193            .streaming_fragment_vnode_mapping
194            .get_or_insert_with(HashMap::new);
195        if mapping.try_insert(fragment_id, vnode_mapping).is_err() {
196            tracing::info!(
197                "Previous batch vnode mapping not found for fragment {fragment_id}, maybe offline scaling with background ddl"
198            );
199        }
200    }
201
202    pub fn update_streaming_fragment_mapping(
203        &self,
204        fragment_id: FragmentId,
205        vnode_mapping: WorkerSlotMapping,
206    ) {
207        let mut guard = self.inner.write().unwrap();
208        let mapping = guard
209            .streaming_fragment_vnode_mapping
210            .get_or_insert_with(HashMap::new);
211        if mapping.insert(fragment_id, vnode_mapping).is_none() {
212            tracing::info!(
213                "Previous vnode mapping not found for fragment {fragment_id}, maybe offline scaling with background ddl"
214            );
215        }
216    }
217
218    pub fn remove_streaming_fragment_mapping(&self, fragment_id: &FragmentId) {
219        let mut guard = self.inner.write().unwrap();
220
221        let res = guard
222            .streaming_fragment_vnode_mapping
223            .as_mut()
224            .and_then(|mapping| mapping.remove(fragment_id));
225        match &res {
226            Some(_) => {}
227            None if fragment_id.is_placeholder() => {
228                // Do nothing for placeholder fragment.
229            }
230            None => {
231                tracing::warn!(%fragment_id, "Streaming vnode mapping not found");
232            }
233        };
234    }
235
236    /// Returns fragment's vnode mapping for serving.
237    fn serving_fragment_mapping(&self, fragment_id: FragmentId) -> Result<WorkerSlotMapping> {
238        self.inner
239            .read()
240            .unwrap()
241            .get_serving_fragment_mapping(fragment_id)
242            .ok_or_else(|| BatchError::ServingVnodeMappingNotFound(fragment_id))
243    }
244
245    pub fn set_serving_fragment_mapping(&self, mappings: HashMap<FragmentId, WorkerSlotMapping>) {
246        let mut guard = self.inner.write().unwrap();
247        tracing::debug!(
248            "Set serving vnode mapping for fragments {:?}",
249            mappings.keys()
250        );
251        guard.serving_fragment_vnode_mapping = mappings;
252    }
253
254    pub fn upsert_serving_fragment_mapping(
255        &self,
256        mappings: HashMap<FragmentId, WorkerSlotMapping>,
257    ) {
258        let mut guard = self.inner.write().unwrap();
259        tracing::debug!(
260            "Upsert serving vnode mapping for fragments {:?}",
261            mappings.keys()
262        );
263        for (fragment_id, mapping) in mappings {
264            guard
265                .serving_fragment_vnode_mapping
266                .insert(fragment_id, mapping);
267        }
268    }
269
270    pub fn remove_serving_fragment_mapping(&self, fragment_ids: &[FragmentId]) {
271        let mut guard = self.inner.write().unwrap();
272        tracing::debug!(
273            "Delete serving vnode mapping for fragments {:?}",
274            fragment_ids
275        );
276        for fragment_id in fragment_ids {
277            guard.serving_fragment_vnode_mapping.remove(fragment_id);
278        }
279    }
280
281    #[cfg(test)]
282    fn worker_node_mask(&self) -> std::sync::RwLockReadGuard<'_, HashSet<WorkerId>> {
283        self.worker_node_mask.read().unwrap()
284    }
285
286    fn worker_node_mask_snapshot(&self) -> HashSet<WorkerId> {
287        self.worker_node_mask.read().unwrap().clone()
288    }
289
290    pub fn mask_worker_node(&self, worker_node_id: WorkerId, duration: Duration) {
291        tracing::info!(
292            "Mask worker node {} for {:?} temporarily",
293            worker_node_id,
294            duration
295        );
296        let mut worker_node_mask = self.worker_node_mask.write().unwrap();
297        if worker_node_mask.contains(&worker_node_id) {
298            return;
299        }
300        worker_node_mask.insert(worker_node_id);
301        let worker_node_mask_ref = self.worker_node_mask.clone();
302        tokio::spawn(async move {
303            tokio::time::sleep(duration).await;
304            worker_node_mask_ref
305                .write()
306                .unwrap()
307                .remove(&worker_node_id);
308        });
309    }
310
311    pub fn worker_node(&self, worker_id: WorkerId) -> Option<WorkerNode> {
312        self.inner.read().unwrap().worker_node(worker_id)
313    }
314}
315
316impl WorkerNodeManagerInner {
317    fn get_serving_fragment_mapping(&self, fragment_id: FragmentId) -> Option<WorkerSlotMapping> {
318        self.serving_fragment_vnode_mapping
319            .get(&fragment_id)
320            .cloned()
321    }
322
323    fn worker_node(&self, worker_id: WorkerId) -> Option<WorkerNode> {
324        self.worker_nodes.get(&worker_id).cloned()
325    }
326}
327
328/// Selects workers for query according to `enable_barrier_read`
329#[derive(Clone)]
330pub struct WorkerNodeSelector {
331    pub manager: WorkerNodeManagerRef,
332    enable_barrier_read: bool,
333}
334
335impl WorkerNodeSelector {
336    pub fn new(manager: WorkerNodeManagerRef, enable_barrier_read: bool) -> Self {
337        Self {
338            manager,
339            enable_barrier_read,
340        }
341    }
342
343    pub fn worker_node_count(&self) -> usize {
344        if self.enable_barrier_read {
345            self.manager.list_streaming_worker_nodes().len()
346        } else {
347            self.apply_worker_node_mask(self.manager.list_serving_worker_nodes())
348                .len()
349        }
350    }
351
352    pub fn schedule_unit_count(&self) -> usize {
353        let worker_nodes = if self.enable_barrier_read {
354            self.manager.list_streaming_worker_nodes()
355        } else {
356            self.apply_worker_node_mask(self.manager.list_serving_worker_nodes())
357        };
358        worker_nodes
359            .iter()
360            .map(|node| node.compute_node_parallelism())
361            .sum()
362    }
363
364    pub fn fragment_mapping(
365        &self,
366        fragment_id: FragmentId,
367        batch_parallelism: usize,
368    ) -> Result<WorkerSlotMapping> {
369        if self.enable_barrier_read {
370            self.manager.get_streaming_fragment_mapping(&fragment_id)
371        } else {
372            let mapping = (self.manager.serving_fragment_mapping(fragment_id)).or_else(|_| {
373                tracing::warn!(
374                    %fragment_id,
375                    "Serving fragment mapping not found, fall back to streaming one."
376                );
377                self.manager.get_streaming_fragment_mapping(&fragment_id)
378            })?;
379            let workers = self.apply_worker_node_mask(self.manager.list_serving_worker_nodes());
380
381            // Filter out unavailable workers.
382            if workers.is_empty() {
383                Err(BatchError::EmptyWorkerNodes)
384            } else {
385                let worker_ids = workers
386                    .iter()
387                    .map(|worker| worker.id)
388                    .collect::<HashSet<_>>();
389                let mapping_uses_only_available_workers = mapping
390                    .iter_unique()
391                    .all(|slot| worker_ids.contains(&slot.worker_id()));
392                if mapping_uses_only_available_workers {
393                    Ok(mapping)
394                } else {
395                    // If it's a singleton, set max_parallelism=1 for place_vnode.
396                    // Otherwise, cap re-placement by the query's effective batch parallelism.
397                    let max_parallelism =
398                        mapping.to_single().map(|_| 1).or(Some(batch_parallelism));
399                    let masked_mapping =
400                        place_vnode(Some(&mapping), &workers, max_parallelism, mapping.len())
401                            .ok_or_else(|| BatchError::EmptyWorkerNodes)?;
402                    Ok(masked_mapping)
403                }
404            }
405        }
406    }
407
408    pub fn next_random_worker(&self) -> Result<WorkerNode> {
409        let worker_nodes = if self.enable_barrier_read {
410            self.manager.list_streaming_worker_nodes()
411        } else {
412            self.apply_worker_node_mask(self.manager.list_serving_worker_nodes())
413        };
414        worker_nodes
415            .choose(&mut rand::rng())
416            .ok_or_else(|| BatchError::EmptyWorkerNodes)
417            .map(|w| (*w).clone())
418    }
419
420    fn is_current_version_worker(worker: &WorkerNode, current_rw_version: &str) -> bool {
421        worker
422            .resource
423            .as_ref()
424            .is_some_and(|resource| resource.rw_version == current_rw_version)
425    }
426
427    fn apply_worker_node_mask(&self, origin: Vec<WorkerNode>) -> Vec<WorkerNode> {
428        let current_rw_version = current_rw_version();
429        let workers_with_current_version = origin
430            .into_iter()
431            .filter(|worker| Self::is_current_version_worker(worker, &current_rw_version))
432            .collect::<Vec<_>>();
433        let worker_node_mask = self.manager.worker_node_mask_snapshot();
434        Self::apply_worker_node_mask_inner(workers_with_current_version, &worker_node_mask)
435    }
436
437    fn apply_worker_node_mask_inner(
438        origin: Vec<WorkerNode>,
439        worker_node_mask: &HashSet<WorkerId>,
440    ) -> Vec<WorkerNode> {
441        if origin.iter().all(|w| worker_node_mask.contains(&w.id)) {
442            return origin;
443        }
444        origin
445            .into_iter()
446            .filter(|w| !worker_node_mask.contains(&w.id))
447            .collect()
448    }
449}
450
451#[cfg(test)]
452mod tests {
453    use std::sync::Arc;
454
455    use itertools::Itertools;
456    use risingwave_common::RW_VERSION;
457    use risingwave_common::util::addr::HostAddr;
458    use risingwave_pb::common::worker_node;
459    use risingwave_pb::common::worker_node::Property;
460
461    use super::*;
462
463    #[test]
464    fn test_worker_node_manager() {
465        let manager = WorkerNodeManager::mock(vec![]);
466        assert_eq!(manager.list_serving_worker_nodes().len(), 0);
467        assert_eq!(manager.list_streaming_worker_nodes().len(), 0);
468        assert_eq!(manager.list_compute_nodes(), vec![]);
469
470        let worker_nodes = vec![
471            WorkerNode {
472                id: 1.into(),
473                r#type: WorkerType::ComputeNode as i32,
474                host: Some(HostAddr::try_from("127.0.0.1:1234").unwrap().to_protobuf()),
475                state: worker_node::State::Running as i32,
476                property: Some(Property {
477                    is_serving: true,
478                    is_streaming: true,
479                    ..Default::default()
480                }),
481                transactional_id: Some(1),
482                ..Default::default()
483            },
484            WorkerNode {
485                id: 2.into(),
486                r#type: WorkerType::ComputeNode as i32,
487                host: Some(HostAddr::try_from("127.0.0.1:1235").unwrap().to_protobuf()),
488                state: worker_node::State::Running as i32,
489                property: Some(Property {
490                    is_serving: true,
491                    is_streaming: false,
492                    ..Default::default()
493                }),
494                transactional_id: Some(2),
495                ..Default::default()
496            },
497        ];
498        worker_nodes
499            .iter()
500            .for_each(|w| manager.add_worker_node(w.clone()));
501        assert_eq!(manager.list_serving_worker_nodes().len(), 2);
502        assert_eq!(manager.list_streaming_worker_nodes().len(), 1);
503        assert_eq!(
504            manager
505                .list_compute_nodes()
506                .into_iter()
507                .sorted_by_key(|w| w.id)
508                .collect_vec(),
509            worker_nodes
510        );
511
512        manager.remove_worker_node(worker_nodes[0].clone());
513        assert_eq!(manager.list_serving_worker_nodes().len(), 1);
514        assert_eq!(manager.list_streaming_worker_nodes().len(), 0);
515        assert_eq!(
516            manager
517                .list_compute_nodes()
518                .into_iter()
519                .sorted_by_key(|w| w.id)
520                .collect_vec(),
521            worker_nodes.as_slice()[1..].to_vec()
522        );
523    }
524
525    fn serving_worker(id: u32, rw_version: Option<&str>) -> WorkerNode {
526        WorkerNode {
527            id: id.into(),
528            r#type: WorkerType::ComputeNode as i32,
529            host: Some(
530                HostAddr::try_from(format!("127.0.0.1:{}", 1234 + id).as_str())
531                    .unwrap()
532                    .to_protobuf(),
533            ),
534            state: worker_node::State::Running as i32,
535            property: Some(Property {
536                is_serving: true,
537                parallelism: 1,
538                ..Default::default()
539            }),
540            transactional_id: Some(id),
541            resource: rw_version.map(|version| worker_node::Resource {
542                rw_version: version.to_owned(),
543                ..Default::default()
544            }),
545            ..Default::default()
546        }
547    }
548
549    fn fragment_mapping_worker_ids(mapping: &WorkerSlotMapping) -> Vec<WorkerId> {
550        mapping
551            .iter_unique()
552            .map(|worker_slot_id| worker_slot_id.worker_id())
553            .collect_vec()
554    }
555
556    fn worker_slot_mapping(worker_ids: impl IntoIterator<Item = u32>) -> WorkerSlotMapping {
557        WorkerSlotMapping::new_uniform(
558            worker_ids
559                .into_iter()
560                .map(|worker_id| WorkerSlotId::new(worker_id.into(), 0))
561                .collect_vec()
562                .into_iter(),
563            4,
564        )
565    }
566
567    #[test]
568    fn test_fragment_mapping_masks_serving_workers_with_version_mismatch() {
569        let manager = Arc::new(WorkerNodeManager::mock(vec![
570            serving_worker(1, Some(RW_VERSION)),
571            serving_worker(2, Some("different-version")),
572        ]));
573        let selector = WorkerNodeSelector::new(manager.clone(), false);
574        manager
575            .set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));
576
577        let mapping = selector.fragment_mapping(0.into(), 4).unwrap();
578
579        assert_eq!(
580            fragment_mapping_worker_ids(&mapping),
581            vec![WorkerId::new(1)]
582        );
583    }
584
585    #[test]
586    fn test_fragment_mapping_keeps_serving_workers_with_current_version() {
587        let manager = Arc::new(WorkerNodeManager::mock(vec![
588            serving_worker(1, Some(RW_VERSION)),
589            serving_worker(2, Some(RW_VERSION)),
590        ]));
591        let selector = WorkerNodeSelector::new(manager.clone(), false);
592        manager
593            .set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));
594
595        let mapping = selector.fragment_mapping(0.into(), 4).unwrap();
596
597        assert_eq!(
598            fragment_mapping_worker_ids(&mapping),
599            vec![WorkerId::new(1), WorkerId::new(2)]
600        );
601    }
602
603    #[test]
604    fn test_fragment_mapping_masks_serving_workers_without_current_version() {
605        let manager = Arc::new(WorkerNodeManager::mock(vec![
606            serving_worker(1, Some(RW_VERSION)),
607            serving_worker(2, None),
608            serving_worker(3, Some("")),
609        ]));
610        let selector = WorkerNodeSelector::new(manager.clone(), false);
611        manager.set_serving_fragment_mapping(HashMap::from([(
612            0.into(),
613            worker_slot_mapping([1, 2, 3]),
614        )]));
615
616        let mapping = selector.fragment_mapping(0.into(), 4).unwrap();
617
618        assert_eq!(
619            fragment_mapping_worker_ids(&mapping),
620            vec![WorkerId::new(1)]
621        );
622    }
623
624    #[test]
625    fn test_fragment_mapping_errors_when_all_serving_workers_have_version_mismatch() {
626        let manager = Arc::new(WorkerNodeManager::mock(vec![
627            serving_worker(1, Some("different-version")),
628            serving_worker(2, None),
629            serving_worker(3, Some("")),
630        ]));
631        let selector = WorkerNodeSelector::new(manager.clone(), false);
632        manager.set_serving_fragment_mapping(HashMap::from([(
633            0.into(),
634            worker_slot_mapping([1, 2, 3]),
635        )]));
636
637        assert!(matches!(
638            selector.fragment_mapping(0.into(), 4),
639            Err(BatchError::EmptyWorkerNodes)
640        ));
641    }
642
643    #[tokio::test]
644    async fn test_fragment_mapping_temporary_mask_fallback_keeps_version_filter() {
645        let manager = Arc::new(WorkerNodeManager::mock(vec![
646            serving_worker(1, Some(RW_VERSION)),
647            serving_worker(2, Some("different-version")),
648        ]));
649        let selector = WorkerNodeSelector::new(manager.clone(), false);
650        manager
651            .set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));
652        manager.mask_worker_node(1.into(), Duration::from_secs(60));
653
654        let mapping = selector.fragment_mapping(0.into(), 4).unwrap();
655
656        assert_eq!(
657            fragment_mapping_worker_ids(&mapping),
658            vec![WorkerId::new(1)]
659        );
660    }
661
662    #[tokio::test]
663    async fn test_fragment_mapping_version_mask_does_not_clear_temporary_mask() {
664        let manager = Arc::new(WorkerNodeManager::mock(vec![
665            serving_worker(1, Some(RW_VERSION)),
666            serving_worker(2, Some(RW_VERSION)),
667            serving_worker(3, Some("different-version")),
668        ]));
669        let selector = WorkerNodeSelector::new(manager.clone(), false);
670        manager.set_serving_fragment_mapping(HashMap::from([(
671            0.into(),
672            worker_slot_mapping([1, 2, 3]),
673        )]));
674        manager.mask_worker_node(2.into(), Duration::from_secs(60));
675
676        let mapping = selector.fragment_mapping(0.into(), 4).unwrap();
677
678        assert_eq!(
679            fragment_mapping_worker_ids(&mapping),
680            vec![WorkerId::new(1)]
681        );
682        assert!(manager.worker_node_mask().contains(&WorkerId::new(2)));
683    }
684
685    #[test]
686    fn test_fragment_mapping_respects_batch_parallelism_with_masked_workers() {
687        let worker_nodes = vec![
688            WorkerNode {
689                id: 1.into(),
690                r#type: WorkerType::ComputeNode as i32,
691                state: worker_node::State::Running as i32,
692                property: Some(Property {
693                    parallelism: 8,
694                    is_serving: true,
695                    ..Default::default()
696                }),
697                resource: Some(worker_node::Resource {
698                    rw_version: RW_VERSION.to_owned(),
699                    ..Default::default()
700                }),
701                ..Default::default()
702            },
703            WorkerNode {
704                id: 2.into(),
705                r#type: WorkerType::ComputeNode as i32,
706                state: worker_node::State::Running as i32,
707                property: Some(Property {
708                    parallelism: 8,
709                    is_serving: true,
710                    ..Default::default()
711                }),
712                resource: Some(worker_node::Resource {
713                    rw_version: RW_VERSION.to_owned(),
714                    ..Default::default()
715                }),
716                ..Default::default()
717            },
718        ];
719        let manager = Arc::new(WorkerNodeManager::mock(worker_nodes));
720        let selector = WorkerNodeSelector::new(manager.clone(), false);
721        let fragment_id = 1.into();
722        let worker_slot_ids = (0..8)
723            .map(|slot| WorkerSlotId::new(1.into(), slot))
724            .chain((0..8).map(|slot| WorkerSlotId::new(2.into(), slot)))
725            .collect_vec();
726        let mapping = WorkerSlotMapping::build_from_ids(&worker_slot_ids, 32);
727        manager.set_serving_fragment_mapping([(fragment_id, mapping)].into_iter().collect());
728
729        manager.worker_node_mask.write().unwrap().insert(1.into());
730
731        let masked_mapping = selector.fragment_mapping(fragment_id, 3).unwrap();
732        assert_eq!(masked_mapping.iter_unique().count(), 3);
733        assert!(
734            masked_mapping
735                .iter_unique()
736                .all(|worker_slot_id| worker_slot_id.worker_id() == WorkerId::new(2))
737        );
738    }
739}