Support min num routing keys in key-based load balancing policy (#16564)
This commit is contained in:
@@ -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),
|
||||
},
|
||||
},
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
},
|
||||
},
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user