From ff3ddb9d9b47fbf0e171891781788b673a085810 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Tue, 13 Jan 2026 13:38:03 +0800 Subject: [PATCH] Support min num routing keys in key-based load balancing policy (#16564) --- sgl-model-gateway/bindings/python/src/lib.rs | 1 + .../python/src/sglang_router/router_args.py | 6 +- sgl-model-gateway/src/config/types.rs | 4 + sgl-model-gateway/src/main.rs | 3 +- sgl-model-gateway/src/policies/manual.rs | 115 ++++++++++++ .../src/routers/http/pd_router.rs | 14 +- sgl-model-gateway/tests/common/mock_worker.rs | 45 +++-- sgl-model-gateway/tests/common/test_config.rs | 23 ++- .../tests/routing/manual_routing_test.rs | 172 ++++++++++++++++++ 9 files changed, 350 insertions(+), 33 deletions(-) diff --git a/sgl-model-gateway/bindings/python/src/lib.rs b/sgl-model-gateway/bindings/python/src/lib.rs index 33aa56eb4..76f1d9920 100644 --- a/sgl-model-gateway/bindings/python/src/lib.rs +++ b/sgl-model-gateway/bindings/python/src/lib.rs @@ -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), }, }, diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py b/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py index f5b2f70f9..88b48caf3 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py @@ -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", diff --git a/sgl-model-gateway/src/config/types.rs b/sgl-model-gateway/src/config/types.rs index 791258cd4..656461630 100644 --- a/sgl-model-gateway/src/config/types.rs +++ b/sgl-model-gateway/src/config/types.rs @@ -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 diff --git a/sgl-model-gateway/src/main.rs b/sgl-model-gateway/src/main.rs index fb989709a..eac0b197c 100644 --- a/sgl-model-gateway/src/main.rs +++ b/sgl-model-gateway/src/main.rs @@ -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), }, }, diff --git a/sgl-model-gateway/src/policies/manual.rs b/sgl-model-gateway/src/policies/manual.rs index 90a7284da..9555fb509 100644 --- a/sgl-model-gateway/src/policies/manual.rs +++ b/sgl-model-gateway/src/policies/manual.rs @@ -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], healthy_indices: &[usize]) -> us select_min_by(healthy_indices, |idx| workers[idx].load()) } +fn min_group_select(workers: &[Arc], 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 { diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 12b5fd3a9..fba326ed7 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -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) } diff --git a/sgl-model-gateway/tests/common/mock_worker.rs b/sgl-model-gateway/tests/common/mock_worker.rs index c1e9142c5..23d6bb6f5 100755 --- a/sgl-model-gateway/tests/common/mock_worker.rs +++ b/sgl-model-gateway/tests/common/mock_worker.rs @@ -299,10 +299,12 @@ async fn generate_handler( Json(payload): Json, ) -> 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() } } diff --git a/sgl-model-gateway/tests/common/test_config.rs b/sgl-model-gateway/tests/common/test_config.rs index 711fbed03..bc79d2cd1 100644 --- a/sgl-model-gateway/tests/common/test_config.rs +++ b/sgl-model-gateway/tests/common/test_config.rs @@ -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 { + (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 { diff --git a/sgl-model-gateway/tests/routing/manual_routing_test.rs b/sgl-model-gateway/tests/routing/manual_routing_test.rs index c806fef16..3e429771c 100644 --- a/sgl-model-gateway/tests/routing/manual_routing_test.rs +++ b/sgl-model-gateway/tests/routing/manual_routing_test.rs @@ -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 = results.into_iter().collect(); + + let worker_counts: HashMap = + 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 = 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> = 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 = key_to_workers.values().flatten().cloned().collect(); + assert_eq!( + all_workers.len(), + 2, + "Keys should be distributed across both workers" + ); + + ctx.shutdown().await; + } +}