diff --git a/sgl-model-gateway/src/policies/cache_aware.rs b/sgl-model-gateway/src/policies/cache_aware.rs index 07dfe3f68..9aef85a94 100644 --- a/sgl-model-gateway/src/policies/cache_aware.rs +++ b/sgl-model-gateway/src/policies/cache_aware.rs @@ -210,6 +210,63 @@ impl CacheAwarePolicy { ); } } + + fn select_worker_min_load( + &self, + workers: &[Arc], + request_text: &Option<&str>, + healthy_indices: &[usize], + model_id: &str, + // TODO may skip passing this arg (and compute inside function) if this is not bottleneck + max_load: usize, + min_load: usize, + ) -> Option { + // Log load balancing trigger + // TODO may use `&str` + let worker_loads: Vec<(String, usize)> = workers + .iter() + .map(|w| (w.url().to_string(), w.load())) + .collect(); + + // TODO may change text + debug!( + "Load balancing triggered | max: {} | min: {} | workers: {:?}", + max_load, min_load, worker_loads + ); + + RouterMetrics::record_load_balancing_event(); + RouterMetrics::set_load_range(max_load, min_load); + + // Use shortest queue when imbalanced + let min_load_idx = healthy_indices + .iter() + .min_by_key(|&&idx| workers[idx].load()) + .copied()?; + + // Even in imbalanced mode, update the tree to maintain cache state + if let Some(text) = request_text { + // Get the tree reference without locking the entire HashMap + // DashMap only locks the specific shard containing this key + let tree = self.trees.get(model_id).map(|entry| entry.value().clone()); + + if let Some(tree) = tree { + // Now we can work with the tree without holding the HashMap lock + tree.insert(text, workers[min_load_idx].url()); + } else { + debug!( + "Warning: No tree found for model '{}', skipping cache update", + model_id + ); + } + } + + // Increment processed counter + workers[min_load_idx].increment_processed(); + RouterMetrics::record_processed_request(workers[min_load_idx].url()); + RouterMetrics::record_policy_decision(self.name(), workers[min_load_idx].url()); + + Some(min_load_idx) + } } impl LoadBalancingPolicy for CacheAwarePolicy { @@ -245,49 +302,14 @@ impl LoadBalancingPolicy for CacheAwarePolicy { && (max_load as f32) > (min_load as f32 * self.config.balance_rel_threshold); if is_imbalanced { - // Log load balancing trigger - let worker_loads: Vec<(String, usize)> = workers - .iter() - .map(|w| (w.url().to_string(), w.load())) - .collect(); - - debug!( - "Load balancing triggered | max: {} | min: {} | workers: {:?}", - max_load, min_load, worker_loads + return self.select_worker_min_load( + workers, + &request_text, + &healthy_indices, + model_id, + max_load, + min_load, ); - - RouterMetrics::record_load_balancing_event(); - RouterMetrics::set_load_range(max_load, min_load); - - // Use shortest queue when imbalanced - let min_load_idx = healthy_indices - .iter() - .min_by_key(|&&idx| workers[idx].load()) - .copied()?; - - // Even in imbalanced mode, update the tree to maintain cache state - if let Some(text) = request_text { - // Get the tree reference without locking the entire HashMap - // DashMap only locks the specific shard containing this key - let tree = self.trees.get(model_id).map(|entry| entry.value().clone()); - - if let Some(tree) = tree { - // Now we can work with the tree without holding the HashMap lock - tree.insert(text, workers[min_load_idx].url()); - } else { - debug!( - "Warning: No tree found for model '{}', skipping cache update", - model_id - ); - } - } - - // Increment processed counter - workers[min_load_idx].increment_processed(); - RouterMetrics::record_processed_request(workers[min_load_idx].url()); - RouterMetrics::record_policy_decision(self.name(), workers[min_load_idx].url()); - - return Some(min_load_idx); } // Use cache-aware routing when balanced