1use 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
29pub struct WorkerNodeManager {
31 inner: RwLock<WorkerNodeManagerInner>,
32 worker_node_mask: Arc<RwLock<HashSet<WorkerId>>>,
34}
35
36struct WorkerNodeManagerInner {
37 worker_nodes: HashMap<WorkerId, WorkerNode>,
38 streaming_fragment_vnode_mapping: Option<HashMap<FragmentId, WorkerSlotMapping>>,
40 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 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 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 }
230 None => {
231 tracing::warn!(%fragment_id, "Streaming vnode mapping not found");
232 }
233 };
234 }
235
236 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#[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 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 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, ¤t_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}