diff --git a/sgl-model-gateway/src/policies/bucket.rs b/sgl-model-gateway/src/policies/bucket.rs index bdf64020b..5db591e18 100644 --- a/sgl-model-gateway/src/policies/bucket.rs +++ b/sgl-model-gateway/src/policies/bucket.rs @@ -1,1127 +1,1321 @@ -use std::{ - collections::{HashMap, HashSet, VecDeque}, - sync::{Arc, Mutex, RwLock}, - thread, - time::{Duration, SystemTime}, -}; - -use dashmap::DashMap; -use rand::Rng; -use tracing::{debug, error, info, warn}; -use uuid::Uuid; - -use super::{get_healthy_worker_indices, normalize_model_key, BucketConfig, LoadBalancingPolicy}; -use crate::core::Worker; - -#[derive(Debug)] -pub struct BucketPolicy { - config: BucketConfig, - buckets: Arc>>>, - adjustment_handle: Option>, -} - -impl Default for BucketPolicy { - fn default() -> Self { - Self::new() - } -} - -impl Drop for BucketPolicy { - fn drop(&mut self) { - if let Some(handle) = self.adjustment_handle.take() { - drop(handle); - } - } -} - -impl BucketPolicy { - pub fn new() -> Self { - Self::with_config(BucketConfig::default()) - } - - pub fn with_config(config: BucketConfig) -> Self { - let buckets = Arc::new(DashMap::>>::new()); - - let adjustment_handle = { - let buckets_clone = Arc::clone(&buckets); - - let interval_secs = config.bucket_adjust_interval_secs; - - Some(thread::spawn(move || loop { - thread::sleep(Duration::from_secs(interval_secs as u64)); - - for bucket_ref in buckets_clone.iter() { - let model_id = bucket_ref.key(); - let bucket = bucket_ref.value(); - match bucket.write() { - Ok(mut bucket_guard) => { - bucket_guard.adjust_boundary(); - } - Err(e) => { - error!( - "Failed to acquire write lock for bucket {}: {}", - model_id, e - ); - } - } - } - })) - }; - - Self { - config, - buckets, - adjustment_handle, - } - } - - pub fn init_prefill_worker_urls(&self, prefill_workers: &[Arc]) { - // Group workers by model - let mut model_workers: HashMap>> = HashMap::new(); - for worker in prefill_workers { - let model_key = normalize_model_key(worker.model_id()); - model_workers - .entry(model_key.to_string()) - .or_default() - .push(worker); - } - // Initialize bucket for each model - for (model_key, model_workers) in model_workers { - let bucket = self - .buckets - .entry(model_key) - .or_insert_with(|| { - Arc::new(RwLock::new(Bucket::new( - self.config.bucket_adjust_interval_secs * 1000, - ))) - }) - .clone(); - - let worker_urls: Vec = model_workers - .iter() - .map(|worker| worker.url().to_string()) - .collect(); - - let lock_result = bucket.write(); - if let Ok(mut bucket_guard) = lock_result { - bucket_guard.init_prefill_worker_urls(worker_urls); - } else { - error!("Failed to acquire write lock for bucket initialization"); - } - } - } - - pub fn add_prefill_url(&self, worker: &dyn Worker) { - let model_key = normalize_model_key(worker.model_id()); - let bucket = self - .buckets - .entry(model_key.to_string()) - .or_insert_with(|| { - Arc::new(RwLock::new(Bucket::new( - self.config.bucket_adjust_interval_secs * 1000, - ))) - }) - .clone(); - - let lock_result = bucket.write(); - if let Ok(mut bucket_guard) = lock_result { - let worker_url = worker.url().to_string(); - - let prefill_worker_urls_clone = { - let mut prefill_worker_urls = bucket_guard.prefill_worker_urls.lock().unwrap(); - if !prefill_worker_urls.contains(&worker_url) { - prefill_worker_urls.push(worker_url.clone()); - } - let cloned = prefill_worker_urls.clone(); - - let mut chars_per_url = bucket_guard.chars_per_url.lock().unwrap(); - chars_per_url.entry(worker_url.clone()).or_insert(0); - - cloned - }; - - bucket_guard.init_prefill_worker_urls(prefill_worker_urls_clone); - - info!( - "Added worker {} to bucket for model {}", - worker_url, model_key - ); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - - pub fn remove_prefill_url(&self, worker: &dyn Worker) { - let model_key = normalize_model_key(worker.model_id()); - - if let Some(bucket_entry) = self.buckets.get(model_key) { - let bucket = bucket_entry.value(); - let worker_url = worker.url().to_string(); - - let lock_result = bucket.write(); - if let Ok(mut bucket_guard) = lock_result { - let (updated_len, updated_urls) = { - let mut prefill_worker_urls = bucket_guard.prefill_worker_urls.lock().unwrap(); - prefill_worker_urls.retain(|u| u != &worker_url); - let len = prefill_worker_urls.len(); - let urls_clone = prefill_worker_urls.clone(); - - let mut chars_per_url = bucket_guard.chars_per_url.lock().unwrap(); - chars_per_url.remove(&worker_url); - - (len, urls_clone) - }; - - bucket_guard.bucket_cnt = updated_len; - - if updated_len > 0 { - bucket_guard.init_prefill_worker_urls(updated_urls); - } - - info!( - "Removed worker {} from bucket for model {} (remaining workers: {})", - worker_url, model_key, bucket_guard.bucket_cnt - ); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } else { - warn!( - "No bucket found for model {} when trying to remove worker", - model_key - ); - } - } -} - -impl LoadBalancingPolicy for BucketPolicy { - fn select_worker( - &self, - workers: &[Arc], - request_text: Option<&str>, - ) -> Option { - let healthy_indices = get_healthy_worker_indices(workers); - - if healthy_indices.is_empty() { - return None; - } - - let char_count = match request_text { - None => 0, - Some(text) => text.chars().count(), - }; - - // Determine the model for this set of workers (router pre-filters by model) - // All workers should be from the same model - let model_key = normalize_model_key(workers[healthy_indices[0]].model_id()); - - let bucket = self - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - let prefill_url = if let Some(bucket) = bucket { - let (choiced_url, chars_per_url_snapshot) = { - let buc = bucket.read().unwrap(); - let chars_per_url_snapshot = buc.chars_per_url.lock().unwrap().clone(); - let choiced_url = buc.find_boundary(char_count); - (choiced_url, chars_per_url_snapshot) - }; - let max_load = chars_per_url_snapshot.values().copied().max().unwrap_or(0); - let min_load = chars_per_url_snapshot.values().copied().min().unwrap_or(0); - let abs_diff = max_load.saturating_sub(min_load); - let rel_threshold = self.config.balance_rel_threshold * min_load as f32; - let is_imbalanced = - abs_diff > self.config.balance_abs_threshold && max_load as f32 > rel_threshold; - debug!( - "Current PD instance status | is_imbalanced={}", - is_imbalanced - ); - - let mut rng = rand::rng(); - let prefill_url = if is_imbalanced { - debug!("select prefill instance by Load Balance policy"); - let min_url = chars_per_url_snapshot - .iter() - .min_by_key(|(_, &chars)| chars) - .map(|(url, _)| url.clone()) - .unwrap_or_else(|| { - let idx = rng.random_range(0..healthy_indices.len()); - let url = workers[healthy_indices[idx]].url(); - warn!("No URL found, randomly selecting: {}", url); - url.to_string() - }); - min_url - } else { - debug!("select prefill instance by Bucket policy"); - match choiced_url { - Some(url) if !url.is_empty() => url, - _ => { - let idx = rng.random_range(0..healthy_indices.len()); - let selected_url = workers[healthy_indices[idx]].url(); - warn!("Boundary not found, randomly selection: {}", selected_url); - selected_url.to_string() - } - } - }; - - { - let mut buc = bucket.write().unwrap(); - buc.post_process_request(char_count, prefill_url.clone()); - } - - prefill_url - } else { - warn!( - "No bucket found for model {}, randomly selecting healthy worker", - model_key - ); - let mut rng = rand::rng(); - let idx = rng.random_range(0..healthy_indices.len()); - let selected_worker = &workers[healthy_indices[idx]]; - let prefill_url = selected_worker.url().to_string(); - prefill_url - }; - - workers.iter().position(|w| w.url() == prefill_url) - } - - fn name(&self) -> &'static str { - "bucket" - } - - fn needs_request_text(&self) -> bool { - true // Bucket policy needs request text - } - - fn as_any(&self) -> &dyn std::any::Any { - self - } -} - -#[derive(Debug, Clone)] -pub struct Bucket { - l_max: usize, - bucket_cnt: usize, - pub prefill_worker_urls: Arc>>, - load_total: usize, - pub period: usize, - bucket_load: usize, - boundary: Vec, - request_list: VecDeque, - t_req_loads: HashMap, - pub chars_per_url: Arc>>, -} - -#[derive(Debug, Clone)] -pub struct SequencerRequest { - pub id: String, - pub char_cnt: usize, - pub timestamp: SystemTime, - pub prefill_worker_url: String, -} - -#[derive(Debug, Clone)] -pub struct Boundary { - pub url: String, - pub range: [usize; 2], -} - -impl Boundary { - pub fn new(url: String, range: [usize; 2]) -> Self { - Boundary { url, range } - } -} - -impl Bucket { - pub fn new(period: usize) -> Self { - let l_max = 4096; - - let bucket_cnt = 0; - - let load_total = 0; - let bucket_load = 0; - - let t_req_loads = HashMap::new(); - let request_list = VecDeque::new(); - - let initial_map = HashMap::new(); - - let boundary = Vec::new(); - - let prefill_worker_urls = Arc::new(Mutex::new(Vec::new())); - - Bucket { - l_max, - bucket_cnt, - prefill_worker_urls, - load_total, - period, - bucket_load, - boundary, - request_list, - t_req_loads, - chars_per_url: Arc::new(Mutex::new(initial_map)), - } - } - - pub fn init_prefill_worker_urls(&mut self, prefill_worker_urls: Vec) { - let bucket_cnt = prefill_worker_urls.len(); - self.bucket_cnt = bucket_cnt; - let mut urls_lock = self.prefill_worker_urls.lock().unwrap(); - *urls_lock = prefill_worker_urls.clone(); - - let mut chars_lock = self.chars_per_url.lock().unwrap(); - chars_lock.clear(); - - for url in prefill_worker_urls.iter() { - chars_lock.insert(url.clone(), 0); - } - - let worker_cnt = bucket_cnt; - let boundary = if worker_cnt == 0 { - Vec::new() - } else { - let gap = self.l_max / worker_cnt; - self.l_max = usize::MAX; - prefill_worker_urls - .iter() - .enumerate() - .map(|(i, url)| { - let min = i * gap; - let max = if i == worker_cnt - 1 { - self.l_max - } else { - (i + 1) * gap - 1 - }; - Boundary::new(url.clone(), [min, max]) - }) - .collect() - }; - - self.boundary = boundary; - info!("Init boundary:{:?}", self.boundary); - } - - pub fn post_process_request(&mut self, char_cnt: usize, prefill_url: String) { - { - let mut map = self.chars_per_url.lock().unwrap(); - *map.entry(prefill_url.clone()).or_insert(0) += char_cnt; - } - - let now = SystemTime::now(); - let time_window_duration = Duration::from_millis(self.period as u64); - let mut removed_load = 0; - - while let Some(req) = self.request_list.front() { - let expired = match now.duration_since(req.timestamp) { - Ok(duration) => duration > time_window_duration, - Err(_) => true, - }; - - if !expired { - break; - } - - if let Some(removed_req) = self.request_list.pop_front() { - self.t_req_loads.remove(&removed_req.id); - removed_load += removed_req.char_cnt; - - let mut map = self.chars_per_url.lock().unwrap(); - if let Some(count) = map.get_mut(&removed_req.prefill_worker_url) { - *count = count.saturating_sub(removed_req.char_cnt); - } - } - } - - self.load_total = self.load_total.saturating_sub(removed_load); - - let id = Uuid::new_v4().to_string(); - - self.t_req_loads.insert(id.clone(), char_cnt); - - self.request_list.push_back(SequencerRequest { - id, - char_cnt, - timestamp: now, - prefill_worker_url: prefill_url, - }); - - self.load_total = self.load_total.saturating_add(char_cnt); - } - - pub fn find_boundary(&self, char_count: usize) -> Option { - let mut left = 0; - let mut right = self.boundary.len(); - let mut _steps = 0; - - while left < right { - _steps += 1; - let mid = left + (right - left) / 2; - let range = self.boundary[mid].range; - - if char_count < range[0] { - right = mid; - } else if char_count > range[1] { - left = mid + 1; - } else { - return Some(self.boundary[mid].url.clone()); - } - } - None - } - - pub fn get_total_load(&self) -> usize { - self.load_total - } - - fn update_workers_cnt(&mut self) { - let pwu = self.prefill_worker_urls.lock().unwrap(); - self.bucket_cnt = pwu.len(); - - let mut char_map = self.chars_per_url.lock().unwrap(); - let current_urls: HashSet<_> = char_map.keys().cloned().collect(); - let new_urls: HashSet<_> = pwu.iter().cloned().collect(); - - for url in new_urls.difference(¤t_urls) { - char_map.insert(url.clone(), 0); - } - - for url in current_urls.difference(&new_urls) { - if char_map.get(url) == Some(&0) { - char_map.remove(url); - } - } - } - - pub fn adjust_boundary(&mut self) { - if self.t_req_loads.is_empty() { - return; - } - - self.update_workers_cnt(); - let worker_cnt = self.bucket_cnt; - if worker_cnt == 0 { - return; - } - let new_single_bucket_load = self.get_total_load() / worker_cnt; - let old_single_bucket_load = self.bucket_load; - - if new_single_bucket_load <= 2 * old_single_bucket_load - && (old_single_bucket_load <= 2 * new_single_bucket_load && old_single_bucket_load != 0) - { - info!("No need to adjust the bucket boundaries."); - return; - } - - info!("Before adjusting boundary | {:?}", self.boundary); - self.bucket_load = new_single_bucket_load; - let mut new_boundary = Vec::new(); - let mut hist_load: Vec = self.t_req_loads.values().cloned().collect(); - hist_load.sort(); - let mut upper_bound: usize = 0; - let mut last_load_index: usize = 0; - let max_value = usize::MAX; - - let worker_url = { - let guard = self.prefill_worker_urls.lock().unwrap(); - (*guard).clone() - }; - - let mut iter = worker_url.iter().peekable(); - // let mut curr_worker_id = 0; - while let Some(url) = iter.next() { - if last_load_index >= hist_load.len() && iter.peek().is_none() { - new_boundary.push(Boundary::new(url.clone(), [upper_bound, max_value])); - break; - } - let mut load_accumulator = 0; - let mut break_flag = false; - for &load in hist_load[last_load_index..].iter() { - load_accumulator += load; - if load_accumulator >= new_single_bucket_load { - if iter.peek().is_none() { - new_boundary.push(Boundary::new(url.clone(), [upper_bound, max_value])); - break_flag = true; - break; - } - let real_load = upper_bound + new_single_bucket_load; - if load <= upper_bound { - new_boundary.push(Boundary::new(url.clone(), [upper_bound, real_load])); - upper_bound = real_load + 1; - } else { - new_boundary.push(Boundary::new(url.clone(), [upper_bound, load])); - upper_bound = load + 1; - } - last_load_index += 1; - break_flag = true; - break; - } else { - last_load_index += 1; - } - } - if !break_flag { - let mut right_bound_value = upper_bound + new_single_bucket_load; - if iter.peek().is_none() { - right_bound_value = max_value; - new_boundary.push(Boundary::new(url.clone(), [upper_bound, right_bound_value])); - break; - } - new_boundary.push(Boundary::new(url.clone(), [upper_bound, right_bound_value])); - upper_bound = right_bound_value + 1; - } - } - self.boundary = new_boundary; - info!("After adjusting boundary | {:?}", self.boundary); - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::core::{BasicWorkerBuilder, WorkerType}; - - #[tokio::test] - async fn test_load_balancing_conditions() { - // Test 1: Basic load balancing trigger - let config = BucketConfig { - balance_abs_threshold: 32, - balance_rel_threshold: 1.0001, - bucket_adjust_interval_secs: 10, - }; - let policy = BucketPolicy::with_config(config); - let prefill_workers: Vec> = vec![ - Arc::new( - BasicWorkerBuilder::new("http://w1:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w2:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w3:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - ]; - - // Initialize the policy with prefill_workers - policy.init_prefill_worker_urls(&prefill_workers); - - // === Phase S1: Construct bucket boundaries === - // Requests len =33 -> Bucket 1(expected range: 0-33) - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(33))) - .unwrap(); - // Two requests len =34 ->load balancing - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(34))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(34))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(11)).await; - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 33); - assert_eq!(bucket_guard.boundary[1].range[1], 67); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - // === Phase S2: Validate load balancing === - // Three consecutive len=33 requests (Should route to different buckets) - let idx_1 = policy - .select_worker(&prefill_workers, Some(&*"a".repeat(33))) - .unwrap(); - let idx_2 = policy - .select_worker(&prefill_workers, Some(&*"a".repeat(33))) - .unwrap(); - let idx_3 = policy - .select_worker(&prefill_workers, Some(&*"a".repeat(33))) - .unwrap(); - assert_eq!(idx_1, 0, "Should not trigger load balancing"); - assert_ne!(idx_2, idx_3, "Should trigger load balancing"); - assert_ne!(idx_2, 0, "Should trigger load balancing"); - assert_ne!(idx_3, 0, "Should trigger load balancing"); - - // Test 2: Not triggering when absolute threshold not met - let config = BucketConfig { - balance_abs_threshold: 30, - balance_rel_threshold: 2.0, - ..Default::default() - }; - let policy = BucketPolicy::with_config(config); - policy.init_prefill_worker_urls(&prefill_workers); - - // Create load difference below absolute threshold(20 + 8 = 28 < 30) - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); // worker1: 20 - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(8))) - .unwrap(); // worker1: 8 - - // Next request should not use bucket scheduling (no load balancing) - let idx = policy - .select_worker(&prefill_workers, Some("request")) - .unwrap(); - assert_eq!( - idx, 0, - "Should not trigger load balancing when relative threshold not met" - ); - - // Test 3: Not triggering when relative threshold not met - let config = BucketConfig { - balance_abs_threshold: 5, - balance_rel_threshold: 3.0, - ..Default::default() - }; - let policy = BucketPolicy::with_config(config); - policy.init_prefill_worker_urls(&prefill_workers); - - // Create load difference (but relative threshold not met) - // Max/Min ratio = 15/5 = 3.0 - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(15))) - .unwrap(); // worker1: 15 - policy - .select_worker(&prefill_workers, Some("short")) - .unwrap(); // worker2: 5 - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(10))) - .unwrap(); // worker3: 10 - - // Next request should use bucket scheduling (load balancing) - let idx = policy - .select_worker(&prefill_workers, Some("request")) - .unwrap(); - assert_eq!( - idx, 0, - "Should not trigger load balancing when relative threshold not met" - ); - } - - #[tokio::test] - async fn test_adjust_boundary_1() { - // Test configuration: Set high threshold to prevent load balancing policy. - let config = BucketConfig { - balance_abs_threshold: 300, - balance_rel_threshold: 1.0001, - bucket_adjust_interval_secs: 3, - }; - let policy = BucketPolicy::with_config(config); - let prefill_workers: Vec> = vec![ - Arc::new( - BasicWorkerBuilder::new("http://w1:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w2:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w3:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - ]; - - // Initialize the policy with prefill_workers - policy.init_prefill_worker_urls(&prefill_workers); - - // Initial boundary - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 1364); - assert_eq!(bucket_guard.boundary[1].range[1], 2729); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - // ===Phase S1: Initial requests to trigger boundary adjustment === - // Send requests with lengths: [5, 10, 15, 20, 24, 26] (total = 100) - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(5))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(10))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(15))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(24))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(26))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - // Verify boundaries adjusted to: [0, 20], [21, 26], [27, MAX] - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 20); - assert_eq!(bucket_guard.boundary[1].range[1], 26); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - // ===Phase S2: Second set of requests to trigger boundary adjustment === - // Send requests with lengths: [10, 20, 30, 40, 45, 57] (total = 202) - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(10))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(30))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(40))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(45))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(57))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - // Verify boundaries adjusted to: [0, 40], [41, 57], [58, MAX] - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 40); - assert_eq!(bucket_guard.boundary[1].range[1], 57); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - } - - #[tokio::test] - async fn test_adjust_boundary_2() { - let config = BucketConfig { - balance_abs_threshold: 300, - balance_rel_threshold: 1.0001, - bucket_adjust_interval_secs: 3, - }; - let policy = BucketPolicy::with_config(config); - let prefill_workers: Vec> = vec![ - Arc::new( - BasicWorkerBuilder::new("http://w1:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w2:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w3:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - ]; - - // Initialize the policy with prefill_workers - policy.init_prefill_worker_urls(&prefill_workers); - - // Initial boundary - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 1364); - assert_eq!(bucket_guard.boundary[1].range[1], 2729); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - // Send requests with char_count 20 - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 20); - assert_eq!(bucket_guard.boundary[1].range[1], 27); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(7))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 7); - assert_eq!(bucket_guard.boundary[1].range[1], 10); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - } - - #[tokio::test] - async fn test_not_adjust_boundary() { - let config = BucketConfig { - balance_abs_threshold: 300, - balance_rel_threshold: 1.0001, - bucket_adjust_interval_secs: 3, - }; - let policy = BucketPolicy::with_config(config); - let prefill_workers: Vec> = vec![ - Arc::new( - BasicWorkerBuilder::new("http://w1:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w2:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - Arc::new( - BasicWorkerBuilder::new("http://w3:8000") - .worker_type(WorkerType::Regular) - .api_key("test_api_key") - .build(), - ), - ]; - - // Initialize the policy with prefill_workers - policy.init_prefill_worker_urls(&prefill_workers); - - // Initial boundary - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 1364); - assert_eq!(bucket_guard.boundary[1].range[1], 2729); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(5))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(10))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(15))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(24))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(26))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 20); - assert_eq!(bucket_guard.boundary[1].range[1], 26); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(10))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(20))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(30))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(32))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(45))) - .unwrap(); - policy - .select_worker(&prefill_workers, Some(&*"a".repeat(55))) - .unwrap(); - - tokio::time::sleep(Duration::from_secs(4)).await; - { - let model_key = "default"; - - let bucket = policy - .buckets - .get(model_key) - .map(|entry| entry.value().clone()); - if let Some(bucket) = bucket { - let lock_result = bucket.write(); - if let Ok(bucket_guard) = lock_result { - // Expected Boundary: [0, 33] [34, 67] [68, MAX] - assert_eq!(bucket_guard.boundary[0].range[1], 20); - assert_eq!(bucket_guard.boundary[1].range[1], 26); - } else { - error!( - "Failed to acquire write lock for bucket of model {}", - model_key - ); - } - } - } - } -} +use std::{ + collections::{HashMap, HashSet, VecDeque}, + sync::{Arc, Mutex, RwLock}, + thread, + time::{Duration, SystemTime}, +}; + +use dashmap::DashMap; +use rand::Rng; +use tracing::{debug, error, info, warn}; +use uuid::Uuid; + +use super::{ + get_healthy_worker_indices, normalize_model_key, BucketConfig, LoadBalancingPolicy, + SelectWorkerInfo, +}; +use crate::core::Worker; + +#[derive(Debug)] +pub struct BucketPolicy { + config: BucketConfig, + buckets: Arc>>>, + adjustment_handle: Option>, +} + +impl Default for BucketPolicy { + fn default() -> Self { + Self::new() + } +} + +impl Drop for BucketPolicy { + fn drop(&mut self) { + if let Some(handle) = self.adjustment_handle.take() { + drop(handle); + } + } +} + +impl BucketPolicy { + pub fn new() -> Self { + Self::with_config(BucketConfig::default()) + } + + pub fn with_config(config: BucketConfig) -> Self { + let buckets = Arc::new(DashMap::>>::new()); + + let adjustment_handle = { + let buckets_clone = Arc::clone(&buckets); + + let interval_secs = config.bucket_adjust_interval_secs; + + Some(thread::spawn(move || loop { + thread::sleep(Duration::from_secs(interval_secs as u64)); + + for bucket_ref in buckets_clone.iter() { + let model_id = bucket_ref.key(); + let bucket = bucket_ref.value(); + match bucket.write() { + Ok(mut bucket_guard) => { + bucket_guard.adjust_boundary(); + } + Err(e) => { + error!( + "Failed to acquire write lock for bucket {}: {}", + model_id, e + ); + } + } + } + })) + }; + + Self { + config, + buckets, + adjustment_handle, + } + } + + pub fn init_prefill_worker_urls(&self, prefill_workers: &[Arc]) { + // Group workers by model + let mut model_workers: HashMap>> = HashMap::new(); + for worker in prefill_workers { + let model_key = normalize_model_key(worker.model_id()); + model_workers + .entry(model_key.to_string()) + .or_default() + .push(worker); + } + // Initialize bucket for each model + for (model_key, model_workers) in model_workers { + let bucket = self + .buckets + .entry(model_key) + .or_insert_with(|| { + Arc::new(RwLock::new(Bucket::new( + self.config.bucket_adjust_interval_secs * 1000, + ))) + }) + .clone(); + + let worker_urls: Vec = model_workers + .iter() + .map(|worker| worker.url().to_string()) + .collect(); + + let lock_result = bucket.write(); + if let Ok(mut bucket_guard) = lock_result { + bucket_guard.init_prefill_worker_urls(worker_urls); + } else { + error!("Failed to acquire write lock for bucket initialization"); + } + } + } + + pub fn add_prefill_url(&self, worker: &dyn Worker) { + let model_key = normalize_model_key(worker.model_id()); + let bucket = self + .buckets + .entry(model_key.to_string()) + .or_insert_with(|| { + Arc::new(RwLock::new(Bucket::new( + self.config.bucket_adjust_interval_secs * 1000, + ))) + }) + .clone(); + + let lock_result = bucket.write(); + if let Ok(mut bucket_guard) = lock_result { + let worker_url = worker.url().to_string(); + + let prefill_worker_urls_clone = { + let mut prefill_worker_urls = bucket_guard.prefill_worker_urls.lock().unwrap(); + if !prefill_worker_urls.contains(&worker_url) { + prefill_worker_urls.push(worker_url.clone()); + } + let cloned = prefill_worker_urls.clone(); + + let mut chars_per_url = bucket_guard.chars_per_url.lock().unwrap(); + chars_per_url.entry(worker_url.clone()).or_insert(0); + + cloned + }; + + bucket_guard.init_prefill_worker_urls(prefill_worker_urls_clone); + + info!( + "Added worker {} to bucket for model {}", + worker_url, model_key + ); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + + pub fn remove_prefill_url(&self, worker: &dyn Worker) { + let model_key = normalize_model_key(worker.model_id()); + + if let Some(bucket_entry) = self.buckets.get(model_key) { + let bucket = bucket_entry.value(); + let worker_url = worker.url().to_string(); + + let lock_result = bucket.write(); + if let Ok(mut bucket_guard) = lock_result { + let (updated_len, updated_urls) = { + let mut prefill_worker_urls = bucket_guard.prefill_worker_urls.lock().unwrap(); + prefill_worker_urls.retain(|u| u != &worker_url); + let len = prefill_worker_urls.len(); + let urls_clone = prefill_worker_urls.clone(); + + let mut chars_per_url = bucket_guard.chars_per_url.lock().unwrap(); + chars_per_url.remove(&worker_url); + + (len, urls_clone) + }; + + bucket_guard.bucket_cnt = updated_len; + + if updated_len > 0 { + bucket_guard.init_prefill_worker_urls(updated_urls); + } + + info!( + "Removed worker {} from bucket for model {} (remaining workers: {})", + worker_url, model_key, bucket_guard.bucket_cnt + ); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } else { + warn!( + "No bucket found for model {} when trying to remove worker", + model_key + ); + } + } +} + +impl LoadBalancingPolicy for BucketPolicy { + fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { + let healthy_indices = get_healthy_worker_indices(workers); + + if healthy_indices.is_empty() { + return None; + } + + let char_count = match info.request_text { + None => 0, + Some(text) => text.chars().count(), + }; + + // Determine the model for this set of workers (router pre-filters by model) + // All workers should be from the same model + let model_key = normalize_model_key(workers[healthy_indices[0]].model_id()); + + let bucket = self + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + let prefill_url = if let Some(bucket) = bucket { + let (choiced_url, chars_per_url_snapshot) = { + let buc = bucket.read().unwrap(); + let chars_per_url_snapshot = buc.chars_per_url.lock().unwrap().clone(); + let choiced_url = buc.find_boundary(char_count); + (choiced_url, chars_per_url_snapshot) + }; + let max_load = chars_per_url_snapshot.values().copied().max().unwrap_or(0); + let min_load = chars_per_url_snapshot.values().copied().min().unwrap_or(0); + let abs_diff = max_load.saturating_sub(min_load); + let rel_threshold = self.config.balance_rel_threshold * min_load as f32; + let is_imbalanced = + abs_diff > self.config.balance_abs_threshold && max_load as f32 > rel_threshold; + debug!( + "Current PD instance status | is_imbalanced={}", + is_imbalanced + ); + + let mut rng = rand::rng(); + let prefill_url = if is_imbalanced { + debug!("select prefill instance by Load Balance policy"); + let min_url = chars_per_url_snapshot + .iter() + .min_by_key(|(_, &chars)| chars) + .map(|(url, _)| url.clone()) + .unwrap_or_else(|| { + let idx = rng.random_range(0..healthy_indices.len()); + let url = workers[healthy_indices[idx]].url(); + warn!("No URL found, randomly selecting: {}", url); + url.to_string() + }); + min_url + } else { + debug!("select prefill instance by Bucket policy"); + match choiced_url { + Some(url) if !url.is_empty() => url, + _ => { + let idx = rng.random_range(0..healthy_indices.len()); + let selected_url = workers[healthy_indices[idx]].url(); + warn!("Boundary not found, randomly selection: {}", selected_url); + selected_url.to_string() + } + } + }; + + { + let mut buc = bucket.write().unwrap(); + buc.post_process_request(char_count, prefill_url.clone()); + } + + prefill_url + } else { + warn!( + "No bucket found for model {}, randomly selecting healthy worker", + model_key + ); + let mut rng = rand::rng(); + let idx = rng.random_range(0..healthy_indices.len()); + let selected_worker = &workers[healthy_indices[idx]]; + let prefill_url = selected_worker.url().to_string(); + prefill_url + }; + + workers.iter().position(|w| w.url() == prefill_url) + } + + fn name(&self) -> &'static str { + "bucket" + } + + fn needs_request_text(&self) -> bool { + true // Bucket policy needs request text + } + + fn as_any(&self) -> &dyn std::any::Any { + self + } +} + +#[derive(Debug, Clone)] +pub struct Bucket { + l_max: usize, + bucket_cnt: usize, + pub prefill_worker_urls: Arc>>, + load_total: usize, + pub period: usize, + bucket_load: usize, + boundary: Vec, + request_list: VecDeque, + t_req_loads: HashMap, + pub chars_per_url: Arc>>, +} + +#[derive(Debug, Clone)] +pub struct SequencerRequest { + pub id: String, + pub char_cnt: usize, + pub timestamp: SystemTime, + pub prefill_worker_url: String, +} + +#[derive(Debug, Clone)] +pub struct Boundary { + pub url: String, + pub range: [usize; 2], +} + +impl Boundary { + pub fn new(url: String, range: [usize; 2]) -> Self { + Boundary { url, range } + } +} + +impl Bucket { + pub fn new(period: usize) -> Self { + let l_max = 4096; + + let bucket_cnt = 0; + + let load_total = 0; + let bucket_load = 0; + + let t_req_loads = HashMap::new(); + let request_list = VecDeque::new(); + + let initial_map = HashMap::new(); + + let boundary = Vec::new(); + + let prefill_worker_urls = Arc::new(Mutex::new(Vec::new())); + + Bucket { + l_max, + bucket_cnt, + prefill_worker_urls, + load_total, + period, + bucket_load, + boundary, + request_list, + t_req_loads, + chars_per_url: Arc::new(Mutex::new(initial_map)), + } + } + + pub fn init_prefill_worker_urls(&mut self, prefill_worker_urls: Vec) { + let bucket_cnt = prefill_worker_urls.len(); + self.bucket_cnt = bucket_cnt; + let mut urls_lock = self.prefill_worker_urls.lock().unwrap(); + *urls_lock = prefill_worker_urls.clone(); + + let mut chars_lock = self.chars_per_url.lock().unwrap(); + chars_lock.clear(); + + for url in prefill_worker_urls.iter() { + chars_lock.insert(url.clone(), 0); + } + + let worker_cnt = bucket_cnt; + let boundary = if worker_cnt == 0 { + Vec::new() + } else { + let gap = self.l_max / worker_cnt; + self.l_max = usize::MAX; + prefill_worker_urls + .iter() + .enumerate() + .map(|(i, url)| { + let min = i * gap; + let max = if i == worker_cnt - 1 { + self.l_max + } else { + (i + 1) * gap - 1 + }; + Boundary::new(url.clone(), [min, max]) + }) + .collect() + }; + + self.boundary = boundary; + info!("Init boundary:{:?}", self.boundary); + } + + pub fn post_process_request(&mut self, char_cnt: usize, prefill_url: String) { + { + let mut map = self.chars_per_url.lock().unwrap(); + *map.entry(prefill_url.clone()).or_insert(0) += char_cnt; + } + + let now = SystemTime::now(); + let time_window_duration = Duration::from_millis(self.period as u64); + let mut removed_load = 0; + + while let Some(req) = self.request_list.front() { + let expired = match now.duration_since(req.timestamp) { + Ok(duration) => duration > time_window_duration, + Err(_) => true, + }; + + if !expired { + break; + } + + if let Some(removed_req) = self.request_list.pop_front() { + self.t_req_loads.remove(&removed_req.id); + removed_load += removed_req.char_cnt; + + let mut map = self.chars_per_url.lock().unwrap(); + if let Some(count) = map.get_mut(&removed_req.prefill_worker_url) { + *count = count.saturating_sub(removed_req.char_cnt); + } + } + } + + self.load_total = self.load_total.saturating_sub(removed_load); + + let id = Uuid::new_v4().to_string(); + + self.t_req_loads.insert(id.clone(), char_cnt); + + self.request_list.push_back(SequencerRequest { + id, + char_cnt, + timestamp: now, + prefill_worker_url: prefill_url, + }); + + self.load_total = self.load_total.saturating_add(char_cnt); + } + + pub fn find_boundary(&self, char_count: usize) -> Option { + let mut left = 0; + let mut right = self.boundary.len(); + let mut _steps = 0; + + while left < right { + _steps += 1; + let mid = left + (right - left) / 2; + let range = self.boundary[mid].range; + + if char_count < range[0] { + right = mid; + } else if char_count > range[1] { + left = mid + 1; + } else { + return Some(self.boundary[mid].url.clone()); + } + } + None + } + + pub fn get_total_load(&self) -> usize { + self.load_total + } + + fn update_workers_cnt(&mut self) { + let pwu = self.prefill_worker_urls.lock().unwrap(); + self.bucket_cnt = pwu.len(); + + let mut char_map = self.chars_per_url.lock().unwrap(); + let current_urls: HashSet<_> = char_map.keys().cloned().collect(); + let new_urls: HashSet<_> = pwu.iter().cloned().collect(); + + for url in new_urls.difference(¤t_urls) { + char_map.insert(url.clone(), 0); + } + + for url in current_urls.difference(&new_urls) { + if char_map.get(url) == Some(&0) { + char_map.remove(url); + } + } + } + + pub fn adjust_boundary(&mut self) { + if self.t_req_loads.is_empty() { + return; + } + + self.update_workers_cnt(); + let worker_cnt = self.bucket_cnt; + if worker_cnt == 0 { + return; + } + let new_single_bucket_load = self.get_total_load() / worker_cnt; + let old_single_bucket_load = self.bucket_load; + + if new_single_bucket_load <= 2 * old_single_bucket_load + && (old_single_bucket_load <= 2 * new_single_bucket_load && old_single_bucket_load != 0) + { + info!("No need to adjust the bucket boundaries."); + return; + } + + info!("Before adjusting boundary | {:?}", self.boundary); + self.bucket_load = new_single_bucket_load; + let mut new_boundary = Vec::new(); + let mut hist_load: Vec = self.t_req_loads.values().cloned().collect(); + hist_load.sort(); + let mut upper_bound: usize = 0; + let mut last_load_index: usize = 0; + let max_value = usize::MAX; + + let worker_url = { + let guard = self.prefill_worker_urls.lock().unwrap(); + (*guard).clone() + }; + + let mut iter = worker_url.iter().peekable(); + // let mut curr_worker_id = 0; + while let Some(url) = iter.next() { + if last_load_index >= hist_load.len() && iter.peek().is_none() { + new_boundary.push(Boundary::new(url.clone(), [upper_bound, max_value])); + break; + } + let mut load_accumulator = 0; + let mut break_flag = false; + for &load in hist_load[last_load_index..].iter() { + load_accumulator += load; + if load_accumulator >= new_single_bucket_load { + if iter.peek().is_none() { + new_boundary.push(Boundary::new(url.clone(), [upper_bound, max_value])); + break_flag = true; + break; + } + let real_load = upper_bound + new_single_bucket_load; + if load <= upper_bound { + new_boundary.push(Boundary::new(url.clone(), [upper_bound, real_load])); + upper_bound = real_load + 1; + } else { + new_boundary.push(Boundary::new(url.clone(), [upper_bound, load])); + upper_bound = load + 1; + } + last_load_index += 1; + break_flag = true; + break; + } else { + last_load_index += 1; + } + } + if !break_flag { + let mut right_bound_value = upper_bound + new_single_bucket_load; + if iter.peek().is_none() { + right_bound_value = max_value; + new_boundary.push(Boundary::new(url.clone(), [upper_bound, right_bound_value])); + break; + } + new_boundary.push(Boundary::new(url.clone(), [upper_bound, right_bound_value])); + upper_bound = right_bound_value + 1; + } + } + self.boundary = new_boundary; + info!("After adjusting boundary | {:?}", self.boundary); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::core::{BasicWorkerBuilder, WorkerType}; + + #[tokio::test] + async fn test_load_balancing_conditions() { + // Test 1: Basic load balancing trigger + let config = BucketConfig { + balance_abs_threshold: 32, + balance_rel_threshold: 1.0001, + bucket_adjust_interval_secs: 10, + }; + let policy = BucketPolicy::with_config(config); + let prefill_workers: Vec> = vec![ + Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w2:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w3:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + ]; + + // Initialize the policy with prefill_workers + policy.init_prefill_worker_urls(&prefill_workers); + + // === Phase S1: Construct bucket boundaries === + // Requests len =33 -> Bucket 1(expected range: 0-33) + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(33)), + }, + ) + .unwrap(); + // Two requests len =34 ->load balancing + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(34)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(34)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(11)).await; + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 33); + assert_eq!(bucket_guard.boundary[1].range[1], 67); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + // === Phase S2: Validate load balancing === + // Three consecutive len=33 requests (Should route to different buckets) + let idx_1 = policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(33)), + }, + ) + .unwrap(); + let idx_2 = policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(33)), + }, + ) + .unwrap(); + let idx_3 = policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(33)), + }, + ) + .unwrap(); + assert_eq!(idx_1, 0, "Should not trigger load balancing"); + assert_ne!(idx_2, idx_3, "Should trigger load balancing"); + assert_ne!(idx_2, 0, "Should trigger load balancing"); + assert_ne!(idx_3, 0, "Should trigger load balancing"); + + // Test 2: Not triggering when absolute threshold not met + let config = BucketConfig { + balance_abs_threshold: 30, + balance_rel_threshold: 2.0, + ..Default::default() + }; + let policy = BucketPolicy::with_config(config); + policy.init_prefill_worker_urls(&prefill_workers); + + // Create load difference below absolute threshold(20 + 8 = 28 < 30) + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); // worker1: 20 + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(8)), + }, + ) + .unwrap(); // worker1: 8 + + // Next request should not use bucket scheduling (no load balancing) + let idx = policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some("request"), + }, + ) + .unwrap(); + assert_eq!( + idx, 0, + "Should not trigger load balancing when relative threshold not met" + ); + + // Test 3: Not triggering when relative threshold not met + let config = BucketConfig { + balance_abs_threshold: 5, + balance_rel_threshold: 3.0, + ..Default::default() + }; + let policy = BucketPolicy::with_config(config); + policy.init_prefill_worker_urls(&prefill_workers); + + // Create load difference (but relative threshold not met) + // Max/Min ratio = 15/5 = 3.0 + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(15)), + }, + ) + .unwrap(); // worker1: 15 + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some("short"), + }, + ) + .unwrap(); // worker2: 5 + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(10)), + }, + ) + .unwrap(); // worker3: 10 + + // Next request should use bucket scheduling (load balancing) + let idx = policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some("request"), + }, + ) + .unwrap(); + assert_eq!( + idx, 0, + "Should not trigger load balancing when relative threshold not met" + ); + } + + #[tokio::test] + async fn test_adjust_boundary_1() { + // Test configuration: Set high threshold to prevent load balancing policy. + let config = BucketConfig { + balance_abs_threshold: 300, + balance_rel_threshold: 1.0001, + bucket_adjust_interval_secs: 3, + }; + let policy = BucketPolicy::with_config(config); + let prefill_workers: Vec> = vec![ + Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w2:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w3:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + ]; + + // Initialize the policy with prefill_workers + policy.init_prefill_worker_urls(&prefill_workers); + + // Initial boundary + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 1364); + assert_eq!(bucket_guard.boundary[1].range[1], 2729); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + // ===Phase S1: Initial requests to trigger boundary adjustment === + // Send requests with lengths: [5, 10, 15, 20, 24, 26] (total = 100) + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(5)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(10)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(15)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(24)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(26)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + // Verify boundaries adjusted to: [0, 20], [21, 26], [27, MAX] + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 20); + assert_eq!(bucket_guard.boundary[1].range[1], 26); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + // ===Phase S2: Second set of requests to trigger boundary adjustment === + // Send requests with lengths: [10, 20, 30, 40, 45, 57] (total = 202) + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(10)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(30)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(40)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(45)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(57)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + // Verify boundaries adjusted to: [0, 40], [41, 57], [58, MAX] + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 40); + assert_eq!(bucket_guard.boundary[1].range[1], 57); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + } + + #[tokio::test] + async fn test_adjust_boundary_2() { + let config = BucketConfig { + balance_abs_threshold: 300, + balance_rel_threshold: 1.0001, + bucket_adjust_interval_secs: 3, + }; + let policy = BucketPolicy::with_config(config); + let prefill_workers: Vec> = vec![ + Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w2:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w3:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + ]; + + // Initialize the policy with prefill_workers + policy.init_prefill_worker_urls(&prefill_workers); + + // Initial boundary + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 1364); + assert_eq!(bucket_guard.boundary[1].range[1], 2729); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + // Send requests with char_count 20 + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 20); + assert_eq!(bucket_guard.boundary[1].range[1], 27); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(7)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 7); + assert_eq!(bucket_guard.boundary[1].range[1], 10); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + } + + #[tokio::test] + async fn test_not_adjust_boundary() { + let config = BucketConfig { + balance_abs_threshold: 300, + balance_rel_threshold: 1.0001, + bucket_adjust_interval_secs: 3, + }; + let policy = BucketPolicy::with_config(config); + let prefill_workers: Vec> = vec![ + Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w2:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + Arc::new( + BasicWorkerBuilder::new("http://w3:8000") + .worker_type(WorkerType::Regular) + .api_key("test_api_key") + .build(), + ), + ]; + + // Initialize the policy with prefill_workers + policy.init_prefill_worker_urls(&prefill_workers); + + // Initial boundary + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 1364); + assert_eq!(bucket_guard.boundary[1].range[1], 2729); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(5)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(10)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(15)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(24)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(26)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 20); + assert_eq!(bucket_guard.boundary[1].range[1], 26); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(10)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(20)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(30)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(32)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(45)), + }, + ) + .unwrap(); + policy + .select_worker( + &prefill_workers, + &SelectWorkerInfo { + request_text: Some(&*"a".repeat(55)), + }, + ) + .unwrap(); + + tokio::time::sleep(Duration::from_secs(4)).await; + { + let model_key = "default"; + + let bucket = policy + .buckets + .get(model_key) + .map(|entry| entry.value().clone()); + if let Some(bucket) = bucket { + let lock_result = bucket.write(); + if let Ok(bucket_guard) = lock_result { + // Expected Boundary: [0, 33] [34, 67] [68, MAX] + assert_eq!(bucket_guard.boundary[0].range[1], 20); + assert_eq!(bucket_guard.boundary[1].range[1], 26); + } else { + error!( + "Failed to acquire write lock for bucket of model {}", + model_key + ); + } + } + } + } +} diff --git a/sgl-model-gateway/src/policies/cache_aware.rs b/sgl-model-gateway/src/policies/cache_aware.rs index ba528d8e8..30cff7d8a 100644 --- a/sgl-model-gateway/src/policies/cache_aware.rs +++ b/sgl-model-gateway/src/policies/cache_aware.rs @@ -74,7 +74,7 @@ use tracing::debug; use super::{ get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig, - LoadBalancingPolicy, + LoadBalancingPolicy, SelectWorkerInfo, }; use crate::core::Worker; @@ -284,11 +284,8 @@ impl CacheAwarePolicy { } impl LoadBalancingPolicy for CacheAwarePolicy { - fn select_worker( - &self, - workers: &[Arc], - request_text: Option<&str>, - ) -> Option { + fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { + let request_text = info.request_text; let healthy_indices = get_healthy_worker_indices(workers); if healthy_indices.is_empty() { @@ -458,14 +455,35 @@ mod tests { policy.init_workers(&workers); // First request should be distributed - let idx1 = policy.select_worker(&workers, Some("hello world")).unwrap(); + let idx1 = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("hello world"), + }, + ) + .unwrap(); // Same request should go to same worker (cache hit) - let idx2 = policy.select_worker(&workers, Some("hello world")).unwrap(); + let idx2 = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("hello world"), + }, + ) + .unwrap(); assert_eq!(idx1, idx2); // Similar request should also go to same worker - let idx3 = policy.select_worker(&workers, Some("hello")).unwrap(); + let idx3 = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("hello"), + }, + ) + .unwrap(); assert_eq!(idx1, idx3); } @@ -496,8 +514,11 @@ mod tests { policy.init_workers(&workers); // Should select worker2 (lower load) despite cache affinity + let info = SelectWorkerInfo { + request_text: Some("test"), + }; for _ in 0..5 { - let idx = policy.select_worker(&workers, Some("test")).unwrap(); + let idx = policy.select_worker(&workers, &info).unwrap(); assert_eq!(idx, 1); // Should always pick worker2 } } @@ -525,15 +546,32 @@ mod tests { policy.init_workers(&workers); // Route some requests - policy.select_worker(&workers, Some("test1")); - policy.select_worker(&workers, Some("test2")); + policy.select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("test1"), + }, + ); + policy.select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("test2"), + }, + ); // Remove a worker policy.remove_worker_by_url("http://w1:8000"); workers[0].set_healthy(false); // All requests should now go to worker2 - let idx = policy.select_worker(&workers, Some("test1")).unwrap(); + let idx = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("test1"), + }, + ) + .unwrap(); assert_eq!(idx, 1); } } diff --git a/sgl-model-gateway/src/policies/mod.rs b/sgl-model-gateway/src/policies/mod.rs index e8ef5fa95..71482dd3b 100644 --- a/sgl-model-gateway/src/policies/mod.rs +++ b/sgl-model-gateway/src/policies/mod.rs @@ -33,11 +33,11 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug { /// /// This is used for regular routing mode where requests go to a single worker. /// Now uses Arc for better performance and to avoid unnecessary cloning. - fn select_worker( - &self, - workers: &[Arc], - request_text: Option<&str>, - ) -> Option; + /// + /// # Arguments + /// * `workers` - Available workers to select from + /// * `info` - Additional information for routing decisions + fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option; /// Update policy state after request completion /// @@ -135,6 +135,13 @@ pub(crate) fn normalize_model_key(model_id: &str) -> &str { } } +/// Information passed to policy for worker selection +#[derive(Debug, Default, Clone)] +pub struct SelectWorkerInfo<'a> { + /// Request text for cache-aware routing + pub request_text: Option<&'a str>, +} + #[cfg(test)] mod tests { use super::*; diff --git a/sgl-model-gateway/src/policies/power_of_two.rs b/sgl-model-gateway/src/policies/power_of_two.rs index 802a51194..425e5a3d1 100644 --- a/sgl-model-gateway/src/policies/power_of_two.rs +++ b/sgl-model-gateway/src/policies/power_of_two.rs @@ -8,7 +8,7 @@ use std::{ use rand::Rng; use tracing::debug; -use super::{get_healthy_worker_indices, LoadBalancingPolicy}; +use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo}; use crate::core::Worker; /// Power-of-two choices policy @@ -33,7 +33,7 @@ impl LoadBalancingPolicy for PowerOfTwoPolicy { fn select_worker( &self, workers: &[Arc], - _request_text: Option<&str>, + _info: &SelectWorkerInfo, ) -> Option { let healthy_indices = get_healthy_worker_indices(workers); @@ -157,8 +157,9 @@ mod tests { // Run multiple selections let mut selected_counts = [0; 3]; + let info = SelectWorkerInfo::default(); for _ in 0..100 { - if let Some(idx) = policy.select_worker(&workers, None) { + if let Some(idx) = policy.select_worker(&workers, &info) { selected_counts[idx] += 1; } } @@ -192,8 +193,9 @@ mod tests { // Should prefer worker2 with lower cached load let mut w2_selected = 0; + let info = SelectWorkerInfo::default(); for _ in 0..50 { - if let Some(idx) = policy.select_worker(&workers, None) { + if let Some(idx) = policy.select_worker(&workers, &info) { if idx == 1 { w2_selected += 1; } @@ -214,7 +216,10 @@ mod tests { )]; // With single worker, should always select it - assert_eq!(policy.select_worker(&workers, None), Some(0)); + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(0) + ); } #[test] @@ -251,7 +256,7 @@ mod tests { // 5. Run selection let selected_idx = policy - .select_worker(&workers, None) + .select_worker(&workers, &SelectWorkerInfo::default()) .expect("Should select a worker"); // 6. Verify the Fix @@ -307,7 +312,9 @@ mod tests { loads_1.insert("http://b:8000".to_string(), 100_000); policy.update_loads(&loads_1); - let idx_1 = policy.select_worker(&workers_1, None).unwrap(); + let idx_1 = policy + .select_worker(&workers_1, &SelectWorkerInfo::default()) + .unwrap(); assert_eq!( idx_1, 0, "Happy Path Failed: Should select Worker A (fewer tokens) despite higher request count" @@ -326,7 +333,9 @@ mod tests { // http://d:8000 is MISSING policy.update_loads(&loads_2); - let idx_2 = policy.select_worker(&workers_2, None).unwrap(); + let idx_2 = policy + .select_worker(&workers_2, &SelectWorkerInfo::default()) + .unwrap(); assert_eq!(idx_2, 1, "Partial Fail 1 Failed: Should fallback to requests and select Worker B (fewer requests)"); // Scenario 3: Partial Failure (Worker A is missing, Worker B has tokens) @@ -342,7 +351,9 @@ mod tests { loads_3.insert("http://f:8000".to_string(), 1_000); policy.update_loads(&loads_3); - let idx_3 = policy.select_worker(&workers_3, None).unwrap(); + let idx_3 = policy + .select_worker(&workers_3, &SelectWorkerInfo::default()) + .unwrap(); assert_eq!(idx_3, 0, "Partial Fail 2 Failed: Should fallback to requests and select Worker A (fewer requests)"); // Scenario 4: Total Failure (Both missing) @@ -356,7 +367,9 @@ mod tests { let loads_4 = HashMap::new(); policy.update_loads(&loads_4); - let idx_4 = policy.select_worker(&workers_4, None).unwrap(); + let idx_4 = policy + .select_worker(&workers_4, &SelectWorkerInfo::default()) + .unwrap(); assert_eq!( idx_4, 1, "Total Fail Failed: Should select Worker B based on request count" diff --git a/sgl-model-gateway/src/policies/random.rs b/sgl-model-gateway/src/policies/random.rs index 12f0ac1dd..0401377a7 100644 --- a/sgl-model-gateway/src/policies/random.rs +++ b/sgl-model-gateway/src/policies/random.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use rand::Rng; -use super::{get_healthy_worker_indices, LoadBalancingPolicy}; +use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo}; use crate::core::Worker; /// Random selection policy @@ -23,7 +23,7 @@ impl LoadBalancingPolicy for RandomPolicy { fn select_worker( &self, workers: &[Arc], - _request_text: Option<&str>, + _info: &SelectWorkerInfo, ) -> Option { let healthy_indices = get_healthy_worker_indices(workers); @@ -76,7 +76,7 @@ mod tests { let mut counts = HashMap::new(); for _ in 0..100 { - if let Some(idx) = policy.select_worker(&workers, None) { + if let Some(idx) = policy.select_worker(&workers, &SelectWorkerInfo::default()) { *counts.entry(idx).or_insert(0) += 1; } } @@ -107,7 +107,10 @@ mod tests { // Should always select the healthy worker (index 1) for _ in 0..10 { - assert_eq!(policy.select_worker(&workers, None), Some(1)); + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(1) + ); } } @@ -121,6 +124,9 @@ mod tests { )]; workers[0].set_healthy(false); - assert_eq!(policy.select_worker(&workers, None), None); + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + None + ); } } diff --git a/sgl-model-gateway/src/policies/round_robin.rs b/sgl-model-gateway/src/policies/round_robin.rs index 739f18ca6..24dba917f 100644 --- a/sgl-model-gateway/src/policies/round_robin.rs +++ b/sgl-model-gateway/src/policies/round_robin.rs @@ -5,7 +5,7 @@ use std::sync::{ Arc, }; -use super::{get_healthy_worker_indices, LoadBalancingPolicy}; +use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo}; use crate::core::Worker; /// Round-robin selection policy @@ -28,7 +28,7 @@ impl LoadBalancingPolicy for RoundRobinPolicy { fn select_worker( &self, workers: &[Arc], - _request_text: Option<&str>, + _info: &SelectWorkerInfo, ) -> Option { let healthy_indices = get_healthy_worker_indices(workers); @@ -83,11 +83,12 @@ mod tests { ]; // Should select workers in order: 0, 1, 2, 0, 1, 2, ... - assert_eq!(policy.select_worker(&workers, None), Some(0)); - assert_eq!(policy.select_worker(&workers, None), Some(1)); - assert_eq!(policy.select_worker(&workers, None), Some(2)); - assert_eq!(policy.select_worker(&workers, None), Some(0)); - assert_eq!(policy.select_worker(&workers, None), Some(1)); + let info = SelectWorkerInfo::default(); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(1)); + assert_eq!(policy.select_worker(&workers, &info), Some(2)); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(1)); } #[test] @@ -115,10 +116,11 @@ mod tests { workers[1].set_healthy(false); // Should skip unhealthy worker: 0, 2, 0, 2, ... - assert_eq!(policy.select_worker(&workers, None), Some(0)); - assert_eq!(policy.select_worker(&workers, None), Some(2)); - assert_eq!(policy.select_worker(&workers, None), Some(0)); - assert_eq!(policy.select_worker(&workers, None), Some(2)); + let info = SelectWorkerInfo::default(); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(2)); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(2)); } #[test] @@ -138,11 +140,12 @@ mod tests { ]; // Advance the counter - assert_eq!(policy.select_worker(&workers, None), Some(0)); - assert_eq!(policy.select_worker(&workers, None), Some(1)); + let info = SelectWorkerInfo::default(); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(1)); // Reset should start from beginning policy.reset(); - assert_eq!(policy.select_worker(&workers, None), Some(0)); + assert_eq!(policy.select_worker(&workers, &info), Some(0)); } } diff --git a/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs b/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs index 17785f3bf..740a3fb00 100644 --- a/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs +++ b/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs @@ -10,7 +10,7 @@ use super::PipelineStage; use crate::{ core::{ConnectionMode, Worker, WorkerRegistry, WorkerType}, observability::metrics::{metrics_labels, Metrics}, - policies::PolicyRegistry, + policies::{PolicyRegistry, SelectWorkerInfo}, routers::{ error, grpc::context::{RequestContext, WorkerSelection}, @@ -146,7 +146,7 @@ impl WorkerSelectionStage { }; // Select worker using the policy - let idx = policy.select_worker(&available, text)?; + let idx = policy.select_worker(&available, &SelectWorkerInfo { request_text: text })?; let selected = available[idx].clone(); // Record worker selection metric @@ -203,8 +203,9 @@ impl WorkerSelectionStage { None => self.policy_registry.get_default_policy(), }; - let prefill_idx = policy.select_worker(&available_prefill, text)?; - let decode_idx = policy.select_worker(&available_decode, text)?; + let info = SelectWorkerInfo { request_text: text }; + let prefill_idx = policy.select_worker(&available_prefill, &info)?; + let decode_idx = policy.select_worker(&available_decode, &info)?; let model = model_id.unwrap_or("default"); let policy_name = policy.name(); diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 6494814b1..94f1d52de 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -26,7 +26,7 @@ use crate::{ metrics::{bool_to_static_str, metrics_labels, Metrics}, otel_trace::inject_trace_context_http, }, - policies::{LoadBalancingPolicy, PolicyRegistry}, + policies::{LoadBalancingPolicy, PolicyRegistry, SelectWorkerInfo}, protocols::{ chat::{ChatCompletionRequest, ChatMessage, MessageContent}, common::{InputIds, StringOrArray}, @@ -784,7 +784,7 @@ impl PDRouter { } let selected_idx = policy - .select_worker(&available_workers, request_text) + .select_worker(&available_workers, &SelectWorkerInfo { request_text }) .ok_or_else(|| { format!( "Policy {} failed to select a {} worker", diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index 281a54706..89e1c12aa 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -27,7 +27,7 @@ use crate::{ metrics::{bool_to_static_str, metrics_labels, Metrics}, otel_trace::inject_trace_context_http, }, - policies::PolicyRegistry, + policies::{PolicyRegistry, SelectWorkerInfo}, protocols::{ chat::ChatCompletionRequest, classify::ClassifyRequest, @@ -168,7 +168,7 @@ impl Router { None => self.policy_registry.get_default_policy(), }; - let idx = policy.select_worker(&available, text)?; + let idx = policy.select_worker(&available, &SelectWorkerInfo { request_text: text })?; // Record worker selection metric (Layer 3) Metrics::record_worker_selection( diff --git a/sgl-model-gateway/tests/cache_aware_backward_compat_test.rs b/sgl-model-gateway/tests/cache_aware_backward_compat_test.rs index 6e45fc372..6412c4129 100644 --- a/sgl-model-gateway/tests/cache_aware_backward_compat_test.rs +++ b/sgl-model-gateway/tests/cache_aware_backward_compat_test.rs @@ -2,7 +2,7 @@ use std::{collections::HashMap, sync::Arc}; use sgl_model_gateway::{ core::{BasicWorkerBuilder, Worker, WorkerType}, - policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy}, + policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy, SelectWorkerInfo}, }; #[test] @@ -40,7 +40,12 @@ fn test_backward_compatibility_with_empty_model_id() { let workers: Vec> = vec![Arc::new(worker1.clone()), Arc::new(worker2.clone())]; // Select worker - should work without errors - let selected = policy.select_worker(&workers, Some("test request")); + let selected = policy.select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("test request"), + }, + ); assert!(selected.is_some(), "Should select a worker"); // Remove workers - should work without errors @@ -97,12 +102,15 @@ fn test_mixed_model_ids() { let default_workers: Vec> = vec![Arc::new(worker1.clone()), Arc::new(worker3.clone())]; - let selected = policy.select_worker(&default_workers, Some("test request")); + let info = SelectWorkerInfo { + request_text: Some("test request"), + }; + let selected = policy.select_worker(&default_workers, &info); assert!(selected.is_some(), "Should select from default workers"); let llama_workers: Vec> = vec![Arc::new(worker2.clone()), Arc::new(worker4.clone())]; - let selected = policy.select_worker(&llama_workers, Some("test request")); + let selected = policy.select_worker(&llama_workers, &info); assert!(selected.is_some(), "Should select from llama-3 workers"); let all_workers: Vec> = vec![ @@ -111,7 +119,7 @@ fn test_mixed_model_ids() { Arc::new(worker3.clone()), Arc::new(worker4.clone()), ]; - let selected = policy.select_worker(&all_workers, Some("test request")); + let selected = policy.select_worker(&all_workers, &info); assert!(selected.is_some(), "Should select from all workers"); } @@ -144,6 +152,11 @@ fn test_remove_worker_by_url_backward_compat() { policy.remove_worker_by_url("http://worker1:8080"); let workers: Vec> = vec![Arc::new(worker2.clone())]; - let selected = policy.select_worker(&workers, Some("test")); + let selected = policy.select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some("test"), + }, + ); assert_eq!(selected, Some(0), "Should only have worker2 left"); }