Support min num routing keys in key-based load balancing policy (#16564)

This commit is contained in:
fzyzcjy
2026-01-13 13:38:03 +08:00
committed by GitHub
parent 9d3018f484
commit ff3ddb9d9b
9 changed files with 350 additions and 33 deletions

View File

@@ -461,6 +461,7 @@ impl Router {
assignment_mode: match self.assignment_mode.as_str() {
"random" => config::ManualAssignmentMode::Random,
"min_load" => config::ManualAssignmentMode::MinLoad,
"min_group" => config::ManualAssignmentMode::MinGroup,
other => panic!("Unknown assignment mode: {}", other),
},
},

View File

@@ -36,7 +36,7 @@ class RouterArgs:
eviction_interval_secs: int = 60
max_tree_size: int = 2**26
max_idle_secs: int = 4 * 3600
assignment_mode: str = "random"
assignment_mode: str = "random" # Mode for manual policy new routing key assignment
max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches
bucket_adjust_interval_secs: int = 5
dp_aware: bool = False
@@ -315,8 +315,8 @@ class RouterArgs:
f"--{prefix}assignment-mode",
type=str,
default=RouterArgs.assignment_mode,
choices=["random", "min_load"],
help="Mode for assigning new routing keys in manual policy: random (default), min_load (worker with fewest requests)",
choices=["random", "min_load", "min_group"],
help="Mode for assigning new routing keys in manual policy: random (default), min_load (worker with fewest requests), min_group (worker with fewest routing keys)",
)
routing_group.add_argument(
f"--{prefix}max-payload-size",

View File

@@ -360,9 +360,13 @@ impl RoutingMode {
#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ManualAssignmentMode {
/// Random selection (default)
#[default]
Random,
/// Select worker with minimum running requests
MinLoad,
/// Select worker with minimum active routing keys
MinGroup,
}
/// Policy configuration for routing

View File

@@ -172,7 +172,7 @@ struct CliArgs {
max_idle_secs: u64,
/// Assignment mode for manual policy when encountering a new routing key
#[arg(long, default_value = "random", value_parser = ["random", "min_load"], help_heading = "Routing Policy")]
#[arg(long, default_value = "random", value_parser = ["random", "min_load", "min_group"], help_heading = "Routing Policy")]
assignment_mode: String,
/// Number of prefix tokens to use for prefix_hash policy
@@ -716,6 +716,7 @@ impl CliArgs {
assignment_mode: match self.assignment_mode.as_str() {
"random" => ManualAssignmentMode::Random,
"min_load" => ManualAssignmentMode::MinLoad,
"min_group" => ManualAssignmentMode::MinGroup,
other => panic!("Unknown assignment mode: {}", other),
},
},

View File

@@ -154,6 +154,7 @@ impl ManualPolicy {
match self.assignment_mode {
ManualAssignmentMode::Random => random_select(healthy_indices),
ManualAssignmentMode::MinLoad => min_load_select(workers, healthy_indices),
ManualAssignmentMode::MinGroup => min_group_select(workers, healthy_indices),
}
}
@@ -290,6 +291,12 @@ fn min_load_select(workers: &[Arc<dyn Worker>], healthy_indices: &[usize]) -> us
select_min_by(healthy_indices, |idx| workers[idx].load())
}
fn min_group_select(workers: &[Arc<dyn Worker>], healthy_indices: &[usize]) -> usize {
select_min_by(healthy_indices, |idx| {
workers[idx].worker_routing_key_load().value()
})
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
@@ -754,6 +761,76 @@ mod tests {
assert_eq!(policy.routing_map.len(), 0);
}
#[test]
fn test_min_group_select_distributes_evenly() {
let config = ManualConfig {
assignment_mode: ManualAssignmentMode::MinGroup,
..Default::default()
};
let policy = ManualPolicy::with_config(config);
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
for i in 0..9 {
let routing_key = format!("key-{}", i);
let headers = headers_with_routing_key(&routing_key);
let info = SelectWorkerInfo {
headers: Some(&headers),
..Default::default()
};
let (result, branch) = policy.select_worker_impl(&workers, &info);
assert!(result.is_some());
assert_eq!(branch, ExecutionBranch::Vacant);
let selected_idx = result.unwrap();
workers[selected_idx]
.worker_routing_key_load()
.increment(&routing_key);
}
let distribution: HashMap<_, usize> = policy
.routing_map
.iter()
.map(|e| e.candi_worker_urls.first().unwrap().clone())
.fold(HashMap::new(), |mut acc, url| {
*acc.entry(url).or_default() += 1;
acc
});
assert_eq!(distribution.len(), 3, "Should use all 3 workers");
for count in distribution.values() {
assert_eq!(*count, 3, "Each worker should have exactly 3 routing keys");
}
}
#[test]
fn test_min_group_select_prefers_worker_with_fewer_routing_keys() {
let config = ManualConfig {
assignment_mode: ManualAssignmentMode::MinGroup,
..Default::default()
};
let policy = ManualPolicy::with_config(config);
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
workers[0].worker_routing_key_load().increment("existing-1");
workers[0].worker_routing_key_load().increment("existing-2");
workers[1].worker_routing_key_load().increment("existing-3");
assert_eq!(workers[0].worker_routing_key_load().value(), 2);
assert_eq!(workers[1].worker_routing_key_load().value(), 1);
assert_eq!(workers[2].worker_routing_key_load().value(), 0);
let headers = headers_with_routing_key("new-key");
let info = SelectWorkerInfo {
headers: Some(&headers),
..Default::default()
};
let (result, _) = policy.select_worker_impl(&workers, &info);
let selected_idx = result.unwrap();
assert_eq!(selected_idx, 2, "Should select worker with 0 routing keys");
}
#[test]
fn test_min_load_select_prefers_worker_with_fewer_requests() {
let config = ManualConfig {
@@ -782,6 +859,44 @@ mod tests {
assert_eq!(selected_idx, 2, "Should select worker with 0 load");
}
#[test]
fn test_min_group_sticky_after_assignment() {
let config = ManualConfig {
assignment_mode: ManualAssignmentMode::MinGroup,
..Default::default()
};
let policy = ManualPolicy::with_config(config);
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
workers[0].worker_routing_key_load().increment("key-0");
workers[1].worker_routing_key_load().increment("key-1");
workers[1].worker_routing_key_load().increment("key-2");
let headers = headers_with_routing_key("new-key");
let info = SelectWorkerInfo {
headers: Some(&headers),
..Default::default()
};
let (first_result, branch) = policy.select_worker_impl(&workers, &info);
let first_idx = first_result.unwrap();
assert_eq!(branch, ExecutionBranch::Vacant);
assert_eq!(
first_idx, 0,
"Should select worker 0 (has 1 routing key vs 2)"
);
for _ in 0..10 {
let (result, branch) = policy.select_worker_impl(&workers, &info);
assert_eq!(
result,
Some(first_idx),
"Same routing key should route to same worker"
);
assert_eq!(branch, ExecutionBranch::OccupiedHit);
}
}
#[test]
fn test_random_mode_does_not_consider_load() {
let config = ManualConfig {

View File

@@ -877,19 +877,17 @@ impl PDRouter {
let stream = UnboundedReceiverStream::new(rx);
let body = Body::from_stream(stream);
let mut response = Response::new(body);
*response.status_mut() = status;
// Attach load guards to response body for proper RAII lifecycle
// Guards are dropped when response body is consumed or client disconnects
let guards = vec![
WorkerLoadGuard::new(prefill, headers.as_ref()),
WorkerLoadGuard::new(decode, headers.as_ref()),
];
let mut headers = headers.unwrap_or_default();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream"));
*response.headers_mut() = headers;
let mut response = Response::new(body);
*response.status_mut() = status;
let mut response_headers = headers.unwrap_or_default();
response_headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream"));
*response.headers_mut() = response_headers;
AttachedBody::wrap_response(response, guards)
}

View File

@@ -299,10 +299,12 @@ async fn generate_handler(
Json(payload): Json<serde_json::Value>,
) -> Response {
let config = config.read().await;
let worker_id = format!("worker-{}", config.port);
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
[("x-worker-id", worker_id)],
Json(json!({
"error": "Random failure for testing"
})),
@@ -373,28 +375,33 @@ async fn generate_handler(
let stream = stream::iter(events);
Sse::new(stream)
.keep_alive(KeepAlive::default())
(
[("x-worker-id", worker_id)],
Sse::new(stream).keep_alive(KeepAlive::default()),
)
.into_response()
} else {
Json(json!({
"text": "This is a mock response.",
"meta_info": {
"prompt_tokens": 10,
"completion_tokens": 5,
"completion_tokens_wo_jump_forward": 5,
"input_token_logprobs": null,
"output_token_logprobs": null,
"first_token_latency": config.response_delay_ms as f64 / 1000.0,
"time_to_first_token": config.response_delay_ms as f64 / 1000.0,
"time_per_output_token": 0.01,
"finish_reason": {
"type": "stop",
"reason": "length"
(
[("x-worker-id", worker_id)],
Json(json!({
"text": "This is a mock response.",
"meta_info": {
"prompt_tokens": 10,
"completion_tokens": 5,
"completion_tokens_wo_jump_forward": 5,
"input_token_logprobs": null,
"output_token_logprobs": null,
"first_token_latency": config.response_delay_ms as f64 / 1000.0,
"time_to_first_token": config.response_delay_ms as f64 / 1000.0,
"time_per_output_token": 0.01,
"finish_reason": {
"type": "stop",
"reason": "length"
}
}
}
}))
.into_response()
})),
)
.into_response()
}
}

View File

@@ -3,7 +3,9 @@
//! Provides pre-configured RouterConfig and MockWorkerConfig builders
//! for common test scenarios.
use smg::config::{CircuitBreakerConfig, PolicyConfig, RetryConfig, RouterConfig};
use smg::config::{
CircuitBreakerConfig, ManualAssignmentMode, PolicyConfig, RetryConfig, RouterConfig,
};
use super::mock_worker::{HealthStatus, MockWorkerConfig, WorkerType};
@@ -94,12 +96,22 @@ impl TestRouterConfig {
/// Create a manual routing config (for sticky routing tests)
pub fn manual(port: u16) -> RouterConfig {
Self::manual_with_mode(port, ManualAssignmentMode::Random)
}
/// Create a manual routing config with min_group assignment mode
pub fn manual_min_group(port: u16) -> RouterConfig {
Self::manual_with_mode(port, ManualAssignmentMode::MinGroup)
}
/// Create a manual routing config with specified assignment mode
pub fn manual_with_mode(port: u16, assignment_mode: ManualAssignmentMode) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.policy(PolicyConfig::Manual {
eviction_interval_secs: 60,
max_idle_secs: 3600,
assignment_mode: Default::default(),
assignment_mode,
})
.host(defaults::HOST)
.port(port)
@@ -262,6 +274,13 @@ impl TestWorkerConfig {
}
}
/// Create multiple slow workers with sequential ports
pub fn slow_workers(start_port: u16, count: u16, delay_ms: u64) -> Vec<MockWorkerConfig> {
(0..count)
.map(|i| Self::slow(start_port + i, delay_ms))
.collect()
}
/// Create a flaky worker config (for retry/fault tolerance tests)
pub fn flaky(port: u16, fail_rate: f32) -> MockWorkerConfig {
MockWorkerConfig {

View File

@@ -172,3 +172,175 @@ mod manual_routing_tests {
ctx.shutdown().await;
}
}
#[cfg(test)]
mod manual_min_group_tests {
use super::*;
async fn send_request(app: axum::Router, routing_key: &str) -> (String, String) {
let payload = json!({
"text": format!("Request for {}", routing_key),
"stream": false
});
let req = Request::builder()
.method("POST")
.uri("/generate")
.header(CONTENT_TYPE, "application/json")
.header(ROUTING_KEY_HEADER, routing_key)
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let worker_id = resp
.headers()
.get("x-worker-id")
.expect("Response should have x-worker-id header")
.to_str()
.unwrap()
.to_string();
(routing_key.to_string(), worker_id)
}
#[tokio::test]
async fn test_min_group_concurrent_distribution() {
let config = TestRouterConfig::manual_min_group(3910);
let ctx =
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29910, 3, 500))
.await;
let app = ctx.create_app().await;
let mut handles = Vec::new();
for i in 0..9 {
let routing_key = format!("key-{}", i);
let app_clone = app.clone();
let handle = tokio::spawn(async move { send_request(app_clone, &routing_key).await });
handles.push(handle);
}
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
let key_to_worker: HashMap<String, String> = results.into_iter().collect();
let worker_counts: HashMap<String, usize> =
key_to_worker.values().fold(HashMap::new(), |mut acc, w| {
*acc.entry(w.clone()).or_default() += 1;
acc
});
assert_eq!(
worker_counts.len(),
3,
"min_group should distribute keys across all 3 workers, got {:?}",
worker_counts
);
for (worker, count) in &worker_counts {
assert_eq!(
*count, 3,
"Worker {} should have exactly 3 keys, got {}. Distribution: {:?}",
worker, count, key_to_worker
);
}
ctx.shutdown().await;
}
#[tokio::test]
async fn test_min_group_sticky_routing() {
let config = TestRouterConfig::manual_min_group(3911);
let ctx =
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29920, 3, 200))
.await;
let app = ctx.create_app().await;
let routing_key = "sticky-key-123";
let mut handles = Vec::new();
for _ in 0..5 {
let app_clone = app.clone();
let key = routing_key.to_string();
let handle = tokio::spawn(async move { send_request(app_clone, &key).await });
handles.push(handle);
}
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
let workers: Vec<String> = results.into_iter().map(|(_, w)| w).collect();
let unique_workers: HashSet<&String> = workers.iter().collect();
assert_eq!(
unique_workers.len(),
1,
"All requests with same routing key should route to same worker, got {:?}",
unique_workers
);
ctx.shutdown().await;
}
#[tokio::test]
async fn test_min_group_mixed_concurrent_routing() {
let config = TestRouterConfig::manual_min_group(3912);
let ctx =
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29930, 2, 300))
.await;
let app = ctx.create_app().await;
let mut handles = Vec::new();
for i in 0..4 {
let routing_key = format!("key-{}", i);
for _ in 0..3 {
let app_clone = app.clone();
let key = routing_key.clone();
let handle = tokio::spawn(async move { send_request(app_clone, &key).await });
handles.push(handle);
}
}
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
let mut key_to_workers: HashMap<String, HashSet<String>> = HashMap::new();
for (key, worker) in results {
key_to_workers.entry(key).or_default().insert(worker);
}
for (key, workers) in &key_to_workers {
assert_eq!(
workers.len(),
1,
"Key {} should route to exactly one worker (sticky), but got {:?}",
key,
workers
);
}
let all_workers: HashSet<String> = key_to_workers.values().flatten().cloned().collect();
assert_eq!(
all_workers.len(),
2,
"Keys should be distributed across both workers"
);
ctx.shutdown().await;
}
}