Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

40 changes: 40 additions & 0 deletions src/batch/src/rpc/service/madsim_test_utils.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
// Copyright 2026 RisingWave Labs
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

use std::collections::HashMap;
use std::sync::{LazyLock, Mutex};

static BATCH_CREATE_TASK_COUNTS: LazyLock<Mutex<HashMap<String, u64>>> =
LazyLock::new(Default::default);

pub fn record_create_task() {
let hostname = risingwave_common::util::resource_util::hostname();
*BATCH_CREATE_TASK_COUNTS
.lock()
.unwrap()
.entry(hostname)
.or_default() += 1;
}

pub fn create_task_count(node: &str) -> u64 {
*BATCH_CREATE_TASK_COUNTS
.lock()
.unwrap()
.get(node)
.unwrap_or(&0)
}

pub fn reset_create_task_counts() {
BATCH_CREATE_TASK_COUNTS.lock().unwrap().clear();
}
2 changes: 2 additions & 0 deletions src/batch/src/rpc/service/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,4 +13,6 @@
// limitations under the License.

pub mod exchange;
#[cfg(madsim)]
pub mod madsim_test_utils;
pub mod task_service;
3 changes: 3 additions & 0 deletions src/batch/src/rpc/service/task_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,9 @@ impl TaskService for BatchServiceImpl {
&self,
request: Request<CreateTaskRequest>,
) -> Result<Response<Self::CreateTaskStream>, Status> {
#[cfg(madsim)]
crate::rpc::service::madsim_test_utils::record_create_task();

let CreateTaskRequest {
task_id,
plan,
Expand Down
244 changes: 225 additions & 19 deletions src/batch/src/worker_manager/worker_node_manager.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,13 +13,14 @@
// limitations under the License.

use std::collections::{HashMap, HashSet};
use std::sync::{Arc, RwLock, RwLockReadGuard};
use std::sync::{Arc, RwLock};
use std::time::Duration;

use rand::seq::IndexedRandom;
use risingwave_common::bail;
use risingwave_common::hash::{WorkerSlotId, WorkerSlotMapping};
use risingwave_common::id::{FragmentId, WorkerId};
use risingwave_common::util::version::current_rw_version;
use risingwave_common::vnode_mapping::vnode_placement::place_vnode;
use risingwave_pb::common::{WorkerNode, WorkerType};

Expand Down Expand Up @@ -277,10 +278,15 @@ impl WorkerNodeManager {
}
}

fn worker_node_mask(&self) -> RwLockReadGuard<'_, HashSet<WorkerId>> {
#[cfg(test)]
fn worker_node_mask(&self) -> std::sync::RwLockReadGuard<'_, HashSet<WorkerId>> {
self.worker_node_mask.read().unwrap()
}

fn worker_node_mask_snapshot(&self) -> HashSet<WorkerId> {
self.worker_node_mask.read().unwrap().clone()
}

pub fn mask_worker_node(&self, worker_node_id: WorkerId, duration: Duration) {
tracing::info!(
"Mask worker node {} for {:?} temporarily",
Expand Down Expand Up @@ -370,19 +376,31 @@ impl WorkerNodeSelector {
);
self.manager.get_streaming_fragment_mapping(&fragment_id)
})?;
let workers = self.apply_worker_node_mask(self.manager.list_serving_worker_nodes());

// Filter out unavailable workers.
if self.manager.worker_node_mask().is_empty() {
Ok(mapping)
if workers.is_empty() {
Err(BatchError::EmptyWorkerNodes)
} else {
let workers = self.apply_worker_node_mask(self.manager.list_serving_worker_nodes());
// If it's a singleton, set max_parallelism=1 for place_vnode.
// Otherwise, cap re-placement by the query's effective batch parallelism.
let max_parallelism = mapping.to_single().map(|_| 1).or(Some(batch_parallelism));
let masked_mapping =
place_vnode(Some(&mapping), &workers, max_parallelism, mapping.len())
.ok_or_else(|| BatchError::EmptyWorkerNodes)?;
Ok(masked_mapping)
let worker_ids = workers
.iter()
.map(|worker| worker.id)
.collect::<HashSet<_>>();
let mapping_uses_only_available_workers = mapping
.iter_unique()
.all(|slot| worker_ids.contains(&slot.worker_id()));
if mapping_uses_only_available_workers {
Ok(mapping)
} else {
// If it's a singleton, set max_parallelism=1 for place_vnode.
// Otherwise, cap re-placement by the query's effective batch parallelism.
let max_parallelism =
mapping.to_single().map(|_| 1).or(Some(batch_parallelism));
let masked_mapping =
place_vnode(Some(&mapping), &workers, max_parallelism, mapping.len())
.ok_or_else(|| BatchError::EmptyWorkerNodes)?;
Ok(masked_mapping)
}
}
}
}
Expand All @@ -399,29 +417,51 @@ impl WorkerNodeSelector {
.map(|w| (*w).clone())
}

fn is_current_version_worker(worker: &WorkerNode, current_rw_version: &str) -> bool {
worker
.resource
.as_ref()
.is_some_and(|resource| resource.rw_version == current_rw_version)
}

fn apply_worker_node_mask(&self, origin: Vec<WorkerNode>) -> Vec<WorkerNode> {
let mask = self.manager.worker_node_mask();
if origin.iter().all(|w| mask.contains(&w.id)) {
let current_rw_version = current_rw_version();
let workers_with_current_version = origin
.into_iter()
.filter(|worker| Self::is_current_version_worker(worker, &current_rw_version))
.collect::<Vec<_>>();
let worker_node_mask = self.manager.worker_node_mask_snapshot();
Self::apply_worker_node_mask_inner(workers_with_current_version, &worker_node_mask)
}

fn apply_worker_node_mask_inner(
origin: Vec<WorkerNode>,
worker_node_mask: &HashSet<WorkerId>,
) -> Vec<WorkerNode> {
if origin.iter().all(|w| worker_node_mask.contains(&w.id)) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Blocking: this fallback now also bypasses the version check. If every serving worker has a different/empty version, or the only matching-version worker is temporarily masked, all(...) is true and we return the original list, making incompatible workers schedulable again. Please keep version mismatch fail-closed and apply the “all masked” fallback only to the temporary mask within the matching-version subset. Please also cover both cases in tests.

@zwang28 zwang28 Aug 4, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

BTW for a typical serving node (frontend + serving role compute node), at least 1 matched serving node is guaranteed.

return origin;
}
origin
.into_iter()
.filter(|w| !mask.contains(&w.id))
.filter(|w| !worker_node_mask.contains(&w.id))
.collect()
}
}

#[cfg(test)]
mod tests {
use std::sync::Arc;

use itertools::Itertools;
use risingwave_common::RW_VERSION;
use risingwave_common::util::addr::HostAddr;
use risingwave_pb::common::worker_node;
use risingwave_pb::common::worker_node::Property;

use super::*;

#[test]
fn test_worker_node_manager() {
use super::*;

let manager = WorkerNodeManager::mock(vec![]);
assert_eq!(manager.list_serving_worker_nodes().len(), 0);
assert_eq!(manager.list_streaming_worker_nodes().len(), 0);
Expand Down Expand Up @@ -482,10 +522,168 @@ mod tests {
);
}

fn serving_worker(id: u32, rw_version: Option<&str>) -> WorkerNode {
WorkerNode {
id: id.into(),
r#type: WorkerType::ComputeNode as i32,
host: Some(
HostAddr::try_from(format!("127.0.0.1:{}", 1234 + id).as_str())
.unwrap()
.to_protobuf(),
),
state: worker_node::State::Running as i32,
property: Some(Property {
is_serving: true,
parallelism: 1,
..Default::default()
}),
transactional_id: Some(id),
resource: rw_version.map(|version| worker_node::Resource {
rw_version: version.to_owned(),
..Default::default()
}),
..Default::default()
}
}

fn fragment_mapping_worker_ids(mapping: &WorkerSlotMapping) -> Vec<WorkerId> {
mapping
.iter_unique()
.map(|worker_slot_id| worker_slot_id.worker_id())
.collect_vec()
}

fn worker_slot_mapping(worker_ids: impl IntoIterator<Item = u32>) -> WorkerSlotMapping {
WorkerSlotMapping::new_uniform(
worker_ids
.into_iter()
.map(|worker_id| WorkerSlotId::new(worker_id.into(), 0))
.collect_vec()
.into_iter(),
4,
)
}

#[test]
fn test_fragment_mapping_masks_serving_workers_with_version_mismatch() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some(RW_VERSION)),
serving_worker(2, Some("different-version")),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager
.set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));

let mapping = selector.fragment_mapping(0.into(), 4).unwrap();

assert_eq!(
fragment_mapping_worker_ids(&mapping),
vec![WorkerId::new(1)]
);
}

#[test]
fn test_fragment_mapping_respects_batch_parallelism_with_masked_workers() {
use super::*;
fn test_fragment_mapping_keeps_serving_workers_with_current_version() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some(RW_VERSION)),
serving_worker(2, Some(RW_VERSION)),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager
.set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));

let mapping = selector.fragment_mapping(0.into(), 4).unwrap();

assert_eq!(
fragment_mapping_worker_ids(&mapping),
vec![WorkerId::new(1), WorkerId::new(2)]
);
}

#[test]
fn test_fragment_mapping_masks_serving_workers_without_current_version() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some(RW_VERSION)),
serving_worker(2, None),
serving_worker(3, Some("")),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager.set_serving_fragment_mapping(HashMap::from([(
0.into(),
worker_slot_mapping([1, 2, 3]),
)]));

let mapping = selector.fragment_mapping(0.into(), 4).unwrap();

assert_eq!(
fragment_mapping_worker_ids(&mapping),
vec![WorkerId::new(1)]
);
}

#[test]
fn test_fragment_mapping_errors_when_all_serving_workers_have_version_mismatch() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some("different-version")),
serving_worker(2, None),
serving_worker(3, Some("")),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager.set_serving_fragment_mapping(HashMap::from([(
0.into(),
worker_slot_mapping([1, 2, 3]),
)]));

assert!(matches!(
selector.fragment_mapping(0.into(), 4),
Err(BatchError::EmptyWorkerNodes)
));
}

#[tokio::test]
async fn test_fragment_mapping_temporary_mask_fallback_keeps_version_filter() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some(RW_VERSION)),
serving_worker(2, Some("different-version")),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager
.set_serving_fragment_mapping(HashMap::from([(0.into(), worker_slot_mapping([1, 2]))]));
manager.mask_worker_node(1.into(), Duration::from_secs(60));

let mapping = selector.fragment_mapping(0.into(), 4).unwrap();

assert_eq!(
fragment_mapping_worker_ids(&mapping),
vec![WorkerId::new(1)]
);
}

#[tokio::test]
async fn test_fragment_mapping_version_mask_does_not_clear_temporary_mask() {
let manager = Arc::new(WorkerNodeManager::mock(vec![
serving_worker(1, Some(RW_VERSION)),
serving_worker(2, Some(RW_VERSION)),
serving_worker(3, Some("different-version")),
]));
let selector = WorkerNodeSelector::new(manager.clone(), false);
manager.set_serving_fragment_mapping(HashMap::from([(
0.into(),
worker_slot_mapping([1, 2, 3]),
)]));
manager.mask_worker_node(2.into(), Duration::from_secs(60));

let mapping = selector.fragment_mapping(0.into(), 4).unwrap();

assert_eq!(
fragment_mapping_worker_ids(&mapping),
vec![WorkerId::new(1)]
);
assert!(manager.worker_node_mask().contains(&WorkerId::new(2)));
}

#[test]
fn test_fragment_mapping_respects_batch_parallelism_with_masked_workers() {
let worker_nodes = vec![
WorkerNode {
id: 1.into(),
Expand All @@ -496,6 +694,10 @@ mod tests {
is_serving: true,
..Default::default()
}),
resource: Some(worker_node::Resource {
rw_version: RW_VERSION.to_owned(),
..Default::default()
}),
..Default::default()
},
WorkerNode {
Expand All @@ -507,6 +709,10 @@ mod tests {
is_serving: true,
..Default::default()
}),
resource: Some(worker_node::Resource {
rw_version: RW_VERSION.to_owned(),
..Default::default()
}),
..Default::default()
},
];
Expand Down
1 change: 1 addition & 0 deletions src/common/src/util/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ pub mod sort_util;
pub mod stream_graph_visitor;
pub mod tracing;
pub mod value_encoding;
pub mod version;
pub mod worker_util;
pub use tokio_util;
pub mod cluster_limit;
Loading
Loading