diff --git a/sgl-model-gateway/bindings/python/src/lib.rs b/sgl-model-gateway/bindings/python/src/lib.rs index 8fdd72dae..33aa56eb4 100644 --- a/sgl-model-gateway/bindings/python/src/lib.rs +++ b/sgl-model-gateway/bindings/python/src/lib.rs @@ -348,6 +348,7 @@ struct Router { eviction_interval_secs: u64, max_tree_size: usize, max_idle_secs: u64, + assignment_mode: String, max_payload_size: usize, dp_aware: bool, api_key: Option, @@ -457,6 +458,11 @@ impl Router { PolicyType::Manual => ConfigPolicyConfig::Manual { eviction_interval_secs: self.eviction_interval_secs, max_idle_secs: self.max_idle_secs, + assignment_mode: match self.assignment_mode.as_str() { + "random" => config::ManualAssignmentMode::Random, + "min_load" => config::ManualAssignmentMode::MinLoad, + other => panic!("Unknown assignment mode: {}", other), + }, }, PolicyType::ConsistentHashing => ConfigPolicyConfig::ConsistentHashing, PolicyType::PrefixHash => ConfigPolicyConfig::PrefixHash { @@ -640,6 +646,7 @@ impl Router { eviction_interval_secs = 120, max_tree_size = 2usize.pow(26), max_idle_secs = 14400, + assignment_mode = String::from("random"), max_payload_size = 512 * 1024 * 1024, dp_aware = false, api_key = None, @@ -724,6 +731,7 @@ impl Router { eviction_interval_secs: u64, max_tree_size: usize, max_idle_secs: u64, + assignment_mode: String, max_payload_size: usize, dp_aware: bool, api_key: Option, @@ -821,6 +829,7 @@ impl Router { eviction_interval_secs, max_tree_size, max_idle_secs, + assignment_mode, max_payload_size, dp_aware, api_key, 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 7140b71a0..f5b2f70f9 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,6 +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" max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches bucket_adjust_interval_secs: int = 5 dp_aware: bool = False @@ -310,6 +311,13 @@ class RouterArgs: default=RouterArgs.max_idle_secs, help="Maximum idle time in seconds before eviction (for manual policy)", ) + routing_group.add_argument( + 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)", + ) routing_group.add_argument( f"--{prefix}max-payload-size", type=int, diff --git a/sgl-model-gateway/src/config/types.rs b/sgl-model-gateway/src/config/types.rs index 597df7bde..791258cd4 100644 --- a/sgl-model-gateway/src/config/types.rs +++ b/sgl-model-gateway/src/config/types.rs @@ -356,6 +356,15 @@ impl RoutingMode { } } +/// Assignment mode for manual policy when encountering a new routing key +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ManualAssignmentMode { + #[default] + Random, + MinLoad, +} + /// Policy configuration for routing #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] @@ -401,6 +410,9 @@ pub enum PolicyConfig { /// Maximum idle time before eviction (seconds, default: 14400 = 4 hours) #[serde(default = "default_manual_max_idle_secs")] max_idle_secs: u64, + /// Assignment mode for new routing keys (default: random) + #[serde(default)] + assignment_mode: ManualAssignmentMode, }, /// Consistent hashing policy using hash ring for session affinity: diff --git a/sgl-model-gateway/src/main.rs b/sgl-model-gateway/src/main.rs index 94751d2b0..fb989709a 100644 --- a/sgl-model-gateway/src/main.rs +++ b/sgl-model-gateway/src/main.rs @@ -5,8 +5,9 @@ use smg::{ auth::{ApiKeyEntry, ControlPlaneAuthConfig, JwtConfig, Role}, config::{ CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, - HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, PostgresConfig, RedisConfig, - RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, TraceConfig, + HistoryBackend, ManualAssignmentMode, MetricsConfig, OracleConfig, PolicyConfig, + PostgresConfig, RedisConfig, RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, + TraceConfig, }, core::ConnectionMode, observability::{ @@ -170,6 +171,10 @@ struct CliArgs { #[arg(long, default_value_t = 14400, help_heading = "Routing Policy")] 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")] + assignment_mode: String, + /// Number of prefix tokens to use for prefix_hash policy #[arg(long, default_value_t = 256, help_heading = "Routing Policy")] prefix_token_count: usize, @@ -708,6 +713,11 @@ impl CliArgs { "manual" => PolicyConfig::Manual { eviction_interval_secs: self.eviction_interval, max_idle_secs: self.max_idle_secs, + assignment_mode: match self.assignment_mode.as_str() { + "random" => ManualAssignmentMode::Random, + "min_load" => ManualAssignmentMode::MinLoad, + other => panic!("Unknown assignment mode: {}", other), + }, }, _ => PolicyConfig::RoundRobin, } diff --git a/sgl-model-gateway/src/policies/factory.rs b/sgl-model-gateway/src/policies/factory.rs index fddb2bc95..ae3ad845e 100644 --- a/sgl-model-gateway/src/policies/factory.rs +++ b/sgl-model-gateway/src/policies/factory.rs @@ -50,10 +50,12 @@ impl PolicyFactory { PolicyConfig::Manual { eviction_interval_secs, max_idle_secs, + assignment_mode, } => { let config = ManualConfig { eviction_interval_secs: *eviction_interval_secs, max_idle_secs: *max_idle_secs, + assignment_mode: *assignment_mode, }; Arc::new(ManualPolicy::with_config(config)) } @@ -125,6 +127,7 @@ mod tests { let policy = PolicyFactory::create_from_config(&PolicyConfig::Manual { eviction_interval_secs: 60, max_idle_secs: 4 * 3600, + assignment_mode: Default::default(), }); assert_eq!(policy.name(), "manual"); diff --git a/sgl-model-gateway/src/policies/manual.rs b/sgl-model-gateway/src/policies/manual.rs index a0f5c2e18..90a7284da 100644 --- a/sgl-model-gateway/src/policies/manual.rs +++ b/sgl-model-gateway/src/policies/manual.rs @@ -23,7 +23,8 @@ use super::{ get_healthy_worker_indices, utils::PeriodicTask, LoadBalancingPolicy, SelectWorkerInfo, }; use crate::{ - core::Worker, observability::metrics::Metrics, routers::header_utils::extract_routing_key, + config::ManualAssignmentMode, core::Worker, observability::metrics::Metrics, + routers::header_utils::extract_routing_key, }; #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -62,6 +63,7 @@ const MAX_CANDIDATE_WORKERS: usize = 2; pub struct ManualConfig { pub eviction_interval_secs: u64, pub max_idle_secs: u64, + pub assignment_mode: ManualAssignmentMode, } impl Default for ManualConfig { @@ -69,6 +71,7 @@ impl Default for ManualConfig { Self { eviction_interval_secs: 60, max_idle_secs: 4 * 3600, + assignment_mode: ManualAssignmentMode::Random, } } } @@ -88,10 +91,10 @@ impl Node { } } -// TODO may optimize performance #[derive(Debug)] pub struct ManualPolicy { routing_map: Arc>, + assignment_mode: ManualAssignmentMode, _eviction_task: Option, } @@ -142,10 +145,18 @@ impl ManualPolicy { Self { routing_map, + assignment_mode: config.assignment_mode, _eviction_task: eviction_task, } } + fn select_new_worker(&self, workers: &[Arc], healthy_indices: &[usize]) -> usize { + match self.assignment_mode { + ManualAssignmentMode::Random => random_select(healthy_indices), + ManualAssignmentMode::MinLoad => min_load_select(workers, healthy_indices), + } + } + fn select_by_routing_id( &self, workers: &[Arc], @@ -163,13 +174,13 @@ impl ManualPolicy { { (idx, ExecutionBranch::OccupiedHit) } else { - let selected_idx = random_select(healthy_indices); + let selected_idx = self.select_new_worker(workers, healthy_indices); node.push_bounded(workers[selected_idx].url().to_string()); (selected_idx, ExecutionBranch::OccupiedMiss) } } Entry::Vacant(entry) => { - let selected_idx = random_select(healthy_indices); + let selected_idx = self.select_new_worker(workers, healthy_indices); entry.insert(Node { candi_worker_urls: vec![workers[selected_idx].url().to_string()], last_access: Instant::now(), @@ -189,15 +200,13 @@ impl ManualPolicy { return (None, ExecutionBranch::NoHealthyWorkers); } - let routing_id = extract_routing_key(info.headers); - - if let Some(routing_id) = routing_id { + if let Some(routing_id) = extract_routing_key(info.headers) { let (idx, branch) = self.select_by_routing_id(workers, routing_id, &healthy_indices); return (Some(idx), branch); } ( - Some(random_select(&healthy_indices)), + Some(self.select_new_worker(workers, &healthy_indices)), ExecutionBranch::NoRoutingId, ) } @@ -239,13 +248,48 @@ fn find_worker_index_by_url(workers: &[Arc], url: &str) -> Option usize { let mut rng = rand::rng(); let random_idx = rng.random_range(0..healthy_indices.len()); healthy_indices[random_idx] } +fn select_min_by(indices: &[K], get_value: F) -> K +where + K: Copy, + V: Ord, + F: Fn(K) -> V, +{ + let mut min_val: Option = None; + let mut candidates = Vec::new(); + + for &idx in indices { + let val = get_value(idx); + match min_val.as_ref().map(|m| val.cmp(m)) { + None | Some(std::cmp::Ordering::Less) => { + min_val = Some(val); + candidates.clear(); + candidates.push(idx); + } + Some(std::cmp::Ordering::Equal) => { + candidates.push(idx); + } + Some(std::cmp::Ordering::Greater) => {} + } + } + + if candidates.len() == 1 { + candidates[0] + } else { + let mut rng = rand::rng(); + candidates[rng.random_range(0..candidates.len())] + } +} + +fn min_load_select(workers: &[Arc], healthy_indices: &[usize]) -> usize { + select_min_by(healthy_indices, |idx| workers[idx].load()) +} + #[cfg(test)] mod tests { use std::collections::HashMap; @@ -642,6 +686,7 @@ mod tests { let config = ManualConfig { eviction_interval_secs: 0, max_idle_secs: 3600, + assignment_mode: ManualAssignmentMode::Random, }; let policy = ManualPolicy::with_config(config); assert!(policy._eviction_task.is_none()); @@ -690,6 +735,7 @@ mod tests { let config = ManualConfig { eviction_interval_secs: 2, max_idle_secs: 2, + assignment_mode: ManualAssignmentMode::Random, }; let policy = ManualPolicy::with_config(config); let workers = create_workers(&["http://w1:8000", "http://w2:8000"]); @@ -707,4 +753,64 @@ mod tests { assert_eq!(policy.routing_map.len(), 0); } + + #[test] + fn test_min_load_select_prefers_worker_with_fewer_requests() { + let config = ManualConfig { + assignment_mode: ManualAssignmentMode::MinLoad, + ..Default::default() + }; + let policy = ManualPolicy::with_config(config); + let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + + workers[0].increment_load(); + workers[0].increment_load(); + workers[1].increment_load(); + + assert_eq!(workers[0].load(), 2); + assert_eq!(workers[1].load(), 1); + assert_eq!(workers[2].load(), 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 load"); + } + + #[test] + fn test_random_mode_does_not_consider_load() { + let config = ManualConfig { + assignment_mode: ManualAssignmentMode::Random, + ..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-1"); + workers[0].worker_routing_key_load().increment("key-2"); + workers[0].worker_routing_key_load().increment("key-3"); + + let mut selected_worker_0 = false; + for i in 0..50 { + let headers = headers_with_routing_key(&format!("test-{}", i)); + let info = SelectWorkerInfo { + headers: Some(&headers), + ..Default::default() + }; + let (result, _) = policy.select_worker_impl(&workers, &info); + if result == Some(0) { + selected_worker_0 = true; + break; + } + } + assert!( + selected_worker_0, + "Random mode should sometimes select worker 0 despite higher load" + ); + } } diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index 8cfa5b01f..4c9c2824c 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -286,15 +286,14 @@ impl Router { } }; - // Optional load tracking for cache-aware policy - // Get the policy for this model to check if it's cache-aware let policy = match model_id { Some(model) => self.policy_registry.get_policy_or_default(model), None => self.policy_registry.get_default_policy(), }; - let load_guard = - (policy.name() == "cache_aware").then(|| WorkerLoadGuard::new(worker.clone(), headers)); + let load_guard = ["cache_aware", "manual"] + .contains(&policy.name()) + .then(|| WorkerLoadGuard::new(worker.clone(), headers)); // Note: Using borrowed reference avoids heap allocation events::RequestSentEvent { url: worker.url() }.emit(); diff --git a/sgl-model-gateway/tests/common/test_config.rs b/sgl-model-gateway/tests/common/test_config.rs index 97483abc0..711fbed03 100644 --- a/sgl-model-gateway/tests/common/test_config.rs +++ b/sgl-model-gateway/tests/common/test_config.rs @@ -99,6 +99,7 @@ impl TestRouterConfig { .policy(PolicyConfig::Manual { eviction_interval_secs: 60, max_idle_secs: 3600, + assignment_mode: Default::default(), }) .host(defaults::HOST) .port(port)