Tiny refactor select_workers API for future passing more information (#15596)

This commit is contained in:
fzyzcjy
2025-12-25 10:31:10 +08:00
committed by GitHub
parent 38dd4fbb66
commit 1ba897f330
10 changed files with 1463 additions and 1188 deletions
+240 -46
View File
@@ -10,7 +10,10 @@ use rand::Rng;
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use uuid::Uuid; use uuid::Uuid;
use super::{get_healthy_worker_indices, normalize_model_key, BucketConfig, LoadBalancingPolicy}; use super::{
get_healthy_worker_indices, normalize_model_key, BucketConfig, LoadBalancingPolicy,
SelectWorkerInfo,
};
use crate::core::Worker; use crate::core::Worker;
#[derive(Debug)] #[derive(Debug)]
@@ -201,18 +204,14 @@ impl BucketPolicy {
} }
impl LoadBalancingPolicy for BucketPolicy { impl LoadBalancingPolicy for BucketPolicy {
fn select_worker( fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
&self,
workers: &[Arc<dyn Worker>],
request_text: Option<&str>,
) -> Option<usize> {
let healthy_indices = get_healthy_worker_indices(workers); let healthy_indices = get_healthy_worker_indices(workers);
if healthy_indices.is_empty() { if healthy_indices.is_empty() {
return None; return None;
} }
let char_count = match request_text { let char_count = match info.request_text {
None => 0, None => 0,
Some(text) => text.chars().count(), Some(text) => text.chars().count(),
}; };
@@ -622,14 +621,29 @@ mod tests {
// === Phase S1: Construct bucket boundaries === // === Phase S1: Construct bucket boundaries ===
// Requests len =33 -> Bucket 1(expected range: 0-33) // Requests len =33 -> Bucket 1(expected range: 0-33)
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(33))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(33)),
},
)
.unwrap(); .unwrap();
// Two requests len =34 ->load balancing // Two requests len =34 ->load balancing
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(34))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(34)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(34))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(34)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(11)).await; tokio::time::sleep(Duration::from_secs(11)).await;
@@ -657,13 +671,28 @@ mod tests {
// === Phase S2: Validate load balancing === // === Phase S2: Validate load balancing ===
// Three consecutive len=33 requests (Should route to different buckets) // Three consecutive len=33 requests (Should route to different buckets)
let idx_1 = policy let idx_1 = policy
.select_worker(&prefill_workers, Some(&*"a".repeat(33))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(33)),
},
)
.unwrap(); .unwrap();
let idx_2 = policy let idx_2 = policy
.select_worker(&prefill_workers, Some(&*"a".repeat(33))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(33)),
},
)
.unwrap(); .unwrap();
let idx_3 = policy let idx_3 = policy
.select_worker(&prefill_workers, Some(&*"a".repeat(33))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(33)),
},
)
.unwrap(); .unwrap();
assert_eq!(idx_1, 0, "Should not trigger load balancing"); assert_eq!(idx_1, 0, "Should not trigger load balancing");
assert_ne!(idx_2, idx_3, "Should trigger load balancing"); assert_ne!(idx_2, idx_3, "Should trigger load balancing");
@@ -681,15 +710,30 @@ mod tests {
// Create load difference below absolute threshold(20 + 8 = 28 < 30) // Create load difference below absolute threshold(20 + 8 = 28 < 30)
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); // worker1: 20 .unwrap(); // worker1: 20
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(8))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(8)),
},
)
.unwrap(); // worker1: 8 .unwrap(); // worker1: 8
// Next request should not use bucket scheduling (no load balancing) // Next request should not use bucket scheduling (no load balancing)
let idx = policy let idx = policy
.select_worker(&prefill_workers, Some("request")) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some("request"),
},
)
.unwrap(); .unwrap();
assert_eq!( assert_eq!(
idx, 0, idx, 0,
@@ -708,18 +752,38 @@ mod tests {
// Create load difference (but relative threshold not met) // Create load difference (but relative threshold not met)
// Max/Min ratio = 15/5 = 3.0 // Max/Min ratio = 15/5 = 3.0
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(15))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(15)),
},
)
.unwrap(); // worker1: 15 .unwrap(); // worker1: 15
policy policy
.select_worker(&prefill_workers, Some("short")) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some("short"),
},
)
.unwrap(); // worker2: 5 .unwrap(); // worker2: 5
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(10))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(10)),
},
)
.unwrap(); // worker3: 10 .unwrap(); // worker3: 10
// Next request should use bucket scheduling (load balancing) // Next request should use bucket scheduling (load balancing)
let idx = policy let idx = policy
.select_worker(&prefill_workers, Some("request")) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some("request"),
},
)
.unwrap(); .unwrap();
assert_eq!( assert_eq!(
idx, 0, idx, 0,
@@ -786,22 +850,52 @@ mod tests {
// ===Phase S1: Initial requests to trigger boundary adjustment === // ===Phase S1: Initial requests to trigger boundary adjustment ===
// Send requests with lengths: [5, 10, 15, 20, 24, 26] (total = 100) // Send requests with lengths: [5, 10, 15, 20, 24, 26] (total = 100)
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(5))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(5)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(10))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(10)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(15))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(15)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(24))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(24)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(26))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(26)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
@@ -831,22 +925,52 @@ mod tests {
// ===Phase S2: Second set of requests to trigger boundary adjustment === // ===Phase S2: Second set of requests to trigger boundary adjustment ===
// Send requests with lengths: [10, 20, 30, 40, 45, 57] (total = 202) // Send requests with lengths: [10, 20, 30, 40, 45, 57] (total = 202)
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(10))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(10)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(30))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(30)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(40))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(40)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(45))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(45)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(57))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(57)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
@@ -931,7 +1055,12 @@ mod tests {
// Send requests with char_count 20 // Send requests with char_count 20
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
@@ -958,7 +1087,12 @@ mod tests {
} }
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(7))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(7)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
@@ -1041,22 +1175,52 @@ mod tests {
} }
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(5))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(5)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(10))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(10)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(15))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(15)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(24))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(24)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(26))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(26)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
@@ -1083,22 +1247,52 @@ mod tests {
} }
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(10))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(10)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(20))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(20)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(30))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(30)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(32))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(32)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(45))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(45)),
},
)
.unwrap(); .unwrap();
policy policy
.select_worker(&prefill_workers, Some(&*"a".repeat(55))) .select_worker(
&prefill_workers,
&SelectWorkerInfo {
request_text: Some(&*"a".repeat(55)),
},
)
.unwrap(); .unwrap();
tokio::time::sleep(Duration::from_secs(4)).await; tokio::time::sleep(Duration::from_secs(4)).await;
+51 -13
View File
@@ -74,7 +74,7 @@ use tracing::debug;
use super::{ use super::{
get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig, get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig,
LoadBalancingPolicy, LoadBalancingPolicy, SelectWorkerInfo,
}; };
use crate::core::Worker; use crate::core::Worker;
@@ -284,11 +284,8 @@ impl CacheAwarePolicy {
} }
impl LoadBalancingPolicy for CacheAwarePolicy { impl LoadBalancingPolicy for CacheAwarePolicy {
fn select_worker( fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
&self, let request_text = info.request_text;
workers: &[Arc<dyn Worker>],
request_text: Option<&str>,
) -> Option<usize> {
let healthy_indices = get_healthy_worker_indices(workers); let healthy_indices = get_healthy_worker_indices(workers);
if healthy_indices.is_empty() { if healthy_indices.is_empty() {
@@ -458,14 +455,35 @@ mod tests {
policy.init_workers(&workers); policy.init_workers(&workers);
// First request should be distributed // 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) // 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); assert_eq!(idx1, idx2);
// Similar request should also go to same worker // 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); assert_eq!(idx1, idx3);
} }
@@ -496,8 +514,11 @@ mod tests {
policy.init_workers(&workers); policy.init_workers(&workers);
// Should select worker2 (lower load) despite cache affinity // Should select worker2 (lower load) despite cache affinity
let info = SelectWorkerInfo {
request_text: Some("test"),
};
for _ in 0..5 { 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 assert_eq!(idx, 1); // Should always pick worker2
} }
} }
@@ -525,15 +546,32 @@ mod tests {
policy.init_workers(&workers); policy.init_workers(&workers);
// Route some requests // Route some requests
policy.select_worker(&workers, Some("test1")); policy.select_worker(
policy.select_worker(&workers, Some("test2")); &workers,
&SelectWorkerInfo {
request_text: Some("test1"),
},
);
policy.select_worker(
&workers,
&SelectWorkerInfo {
request_text: Some("test2"),
},
);
// Remove a worker // Remove a worker
policy.remove_worker_by_url("http://w1:8000"); policy.remove_worker_by_url("http://w1:8000");
workers[0].set_healthy(false); workers[0].set_healthy(false);
// All requests should now go to worker2 // 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); assert_eq!(idx, 1);
} }
} }
+12 -5
View File
@@ -33,11 +33,11 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug {
/// ///
/// This is used for regular routing mode where requests go to a single worker. /// This is used for regular routing mode where requests go to a single worker.
/// Now uses Arc<dyn Worker> for better performance and to avoid unnecessary cloning. /// Now uses Arc<dyn Worker> for better performance and to avoid unnecessary cloning.
fn select_worker( ///
&self, /// # Arguments
workers: &[Arc<dyn Worker>], /// * `workers` - Available workers to select from
request_text: Option<&str>, /// * `info` - Additional information for routing decisions
) -> Option<usize>; fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize>;
/// Update policy state after request completion /// 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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
+23 -10
View File
@@ -8,7 +8,7 @@ use std::{
use rand::Rng; use rand::Rng;
use tracing::debug; use tracing::debug;
use super::{get_healthy_worker_indices, LoadBalancingPolicy}; use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
use crate::core::Worker; use crate::core::Worker;
/// Power-of-two choices policy /// Power-of-two choices policy
@@ -33,7 +33,7 @@ impl LoadBalancingPolicy for PowerOfTwoPolicy {
fn select_worker( fn select_worker(
&self, &self,
workers: &[Arc<dyn Worker>], workers: &[Arc<dyn Worker>],
_request_text: Option<&str>, _info: &SelectWorkerInfo,
) -> Option<usize> { ) -> Option<usize> {
let healthy_indices = get_healthy_worker_indices(workers); let healthy_indices = get_healthy_worker_indices(workers);
@@ -157,8 +157,9 @@ mod tests {
// Run multiple selections // Run multiple selections
let mut selected_counts = [0; 3]; let mut selected_counts = [0; 3];
let info = SelectWorkerInfo::default();
for _ in 0..100 { 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; selected_counts[idx] += 1;
} }
} }
@@ -192,8 +193,9 @@ mod tests {
// Should prefer worker2 with lower cached load // Should prefer worker2 with lower cached load
let mut w2_selected = 0; let mut w2_selected = 0;
let info = SelectWorkerInfo::default();
for _ in 0..50 { 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 { if idx == 1 {
w2_selected += 1; w2_selected += 1;
} }
@@ -214,7 +216,10 @@ mod tests {
)]; )];
// With single worker, should always select it // 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] #[test]
@@ -251,7 +256,7 @@ mod tests {
// 5. Run selection // 5. Run selection
let selected_idx = policy let selected_idx = policy
.select_worker(&workers, None) .select_worker(&workers, &SelectWorkerInfo::default())
.expect("Should select a worker"); .expect("Should select a worker");
// 6. Verify the Fix // 6. Verify the Fix
@@ -307,7 +312,9 @@ mod tests {
loads_1.insert("http://b:8000".to_string(), 100_000); loads_1.insert("http://b:8000".to_string(), 100_000);
policy.update_loads(&loads_1); 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!( assert_eq!(
idx_1, 0, idx_1, 0,
"Happy Path Failed: Should select Worker A (fewer tokens) despite higher request count" "Happy Path Failed: Should select Worker A (fewer tokens) despite higher request count"
@@ -326,7 +333,9 @@ mod tests {
// http://d:8000 is MISSING // http://d:8000 is MISSING
policy.update_loads(&loads_2); 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)"); 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) // 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); loads_3.insert("http://f:8000".to_string(), 1_000);
policy.update_loads(&loads_3); 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)"); 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) // Scenario 4: Total Failure (Both missing)
@@ -356,7 +367,9 @@ mod tests {
let loads_4 = HashMap::new(); let loads_4 = HashMap::new();
policy.update_loads(&loads_4); 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!( assert_eq!(
idx_4, 1, idx_4, 1,
"Total Fail Failed: Should select Worker B based on request count" "Total Fail Failed: Should select Worker B based on request count"
+11 -5
View File
@@ -4,7 +4,7 @@ use std::sync::Arc;
use rand::Rng; use rand::Rng;
use super::{get_healthy_worker_indices, LoadBalancingPolicy}; use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
use crate::core::Worker; use crate::core::Worker;
/// Random selection policy /// Random selection policy
@@ -23,7 +23,7 @@ impl LoadBalancingPolicy for RandomPolicy {
fn select_worker( fn select_worker(
&self, &self,
workers: &[Arc<dyn Worker>], workers: &[Arc<dyn Worker>],
_request_text: Option<&str>, _info: &SelectWorkerInfo,
) -> Option<usize> { ) -> Option<usize> {
let healthy_indices = get_healthy_worker_indices(workers); let healthy_indices = get_healthy_worker_indices(workers);
@@ -76,7 +76,7 @@ mod tests {
let mut counts = HashMap::new(); let mut counts = HashMap::new();
for _ in 0..100 { 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; *counts.entry(idx).or_insert(0) += 1;
} }
} }
@@ -107,7 +107,10 @@ mod tests {
// Should always select the healthy worker (index 1) // Should always select the healthy worker (index 1)
for _ in 0..10 { 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); workers[0].set_healthy(false);
assert_eq!(policy.select_worker(&workers, None), None); assert_eq!(
policy.select_worker(&workers, &SelectWorkerInfo::default()),
None
);
} }
} }
+17 -14
View File
@@ -5,7 +5,7 @@ use std::sync::{
Arc, Arc,
}; };
use super::{get_healthy_worker_indices, LoadBalancingPolicy}; use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
use crate::core::Worker; use crate::core::Worker;
/// Round-robin selection policy /// Round-robin selection policy
@@ -28,7 +28,7 @@ impl LoadBalancingPolicy for RoundRobinPolicy {
fn select_worker( fn select_worker(
&self, &self,
workers: &[Arc<dyn Worker>], workers: &[Arc<dyn Worker>],
_request_text: Option<&str>, _info: &SelectWorkerInfo,
) -> Option<usize> { ) -> Option<usize> {
let healthy_indices = get_healthy_worker_indices(workers); 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, ... // Should select workers in order: 0, 1, 2, 0, 1, 2, ...
assert_eq!(policy.select_worker(&workers, None), Some(0)); let info = SelectWorkerInfo::default();
assert_eq!(policy.select_worker(&workers, None), Some(1)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
assert_eq!(policy.select_worker(&workers, None), Some(2)); assert_eq!(policy.select_worker(&workers, &info), Some(1));
assert_eq!(policy.select_worker(&workers, None), Some(0)); assert_eq!(policy.select_worker(&workers, &info), Some(2));
assert_eq!(policy.select_worker(&workers, None), Some(1)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
assert_eq!(policy.select_worker(&workers, &info), Some(1));
} }
#[test] #[test]
@@ -115,10 +116,11 @@ mod tests {
workers[1].set_healthy(false); workers[1].set_healthy(false);
// Should skip unhealthy worker: 0, 2, 0, 2, ... // Should skip unhealthy worker: 0, 2, 0, 2, ...
assert_eq!(policy.select_worker(&workers, None), Some(0)); let info = SelectWorkerInfo::default();
assert_eq!(policy.select_worker(&workers, None), Some(2)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
assert_eq!(policy.select_worker(&workers, None), Some(0)); assert_eq!(policy.select_worker(&workers, &info), Some(2));
assert_eq!(policy.select_worker(&workers, None), Some(2)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
assert_eq!(policy.select_worker(&workers, &info), Some(2));
} }
#[test] #[test]
@@ -138,11 +140,12 @@ mod tests {
]; ];
// Advance the counter // Advance the counter
assert_eq!(policy.select_worker(&workers, None), Some(0)); let info = SelectWorkerInfo::default();
assert_eq!(policy.select_worker(&workers, None), Some(1)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
assert_eq!(policy.select_worker(&workers, &info), Some(1));
// Reset should start from beginning // Reset should start from beginning
policy.reset(); policy.reset();
assert_eq!(policy.select_worker(&workers, None), Some(0)); assert_eq!(policy.select_worker(&workers, &info), Some(0));
} }
} }
@@ -10,7 +10,7 @@ use super::PipelineStage;
use crate::{ use crate::{
core::{ConnectionMode, Worker, WorkerRegistry, WorkerType}, core::{ConnectionMode, Worker, WorkerRegistry, WorkerType},
observability::metrics::{metrics_labels, Metrics}, observability::metrics::{metrics_labels, Metrics},
policies::PolicyRegistry, policies::{PolicyRegistry, SelectWorkerInfo},
routers::{ routers::{
error, error,
grpc::context::{RequestContext, WorkerSelection}, grpc::context::{RequestContext, WorkerSelection},
@@ -146,7 +146,7 @@ impl WorkerSelectionStage {
}; };
// Select worker using the policy // 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(); let selected = available[idx].clone();
// Record worker selection metric // Record worker selection metric
@@ -203,8 +203,9 @@ impl WorkerSelectionStage {
None => self.policy_registry.get_default_policy(), None => self.policy_registry.get_default_policy(),
}; };
let prefill_idx = policy.select_worker(&available_prefill, text)?; let info = SelectWorkerInfo { request_text: text };
let decode_idx = policy.select_worker(&available_decode, 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 model = model_id.unwrap_or("default");
let policy_name = policy.name(); let policy_name = policy.name();
@@ -26,7 +26,7 @@ use crate::{
metrics::{bool_to_static_str, metrics_labels, Metrics}, metrics::{bool_to_static_str, metrics_labels, Metrics},
otel_trace::inject_trace_context_http, otel_trace::inject_trace_context_http,
}, },
policies::{LoadBalancingPolicy, PolicyRegistry}, policies::{LoadBalancingPolicy, PolicyRegistry, SelectWorkerInfo},
protocols::{ protocols::{
chat::{ChatCompletionRequest, ChatMessage, MessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
common::{InputIds, StringOrArray}, common::{InputIds, StringOrArray},
@@ -784,7 +784,7 @@ impl PDRouter {
} }
let selected_idx = policy let selected_idx = policy
.select_worker(&available_workers, request_text) .select_worker(&available_workers, &SelectWorkerInfo { request_text })
.ok_or_else(|| { .ok_or_else(|| {
format!( format!(
"Policy {} failed to select a {} worker", "Policy {} failed to select a {} worker",
+2 -2
View File
@@ -27,7 +27,7 @@ use crate::{
metrics::{bool_to_static_str, metrics_labels, Metrics}, metrics::{bool_to_static_str, metrics_labels, Metrics},
otel_trace::inject_trace_context_http, otel_trace::inject_trace_context_http,
}, },
policies::PolicyRegistry, policies::{PolicyRegistry, SelectWorkerInfo},
protocols::{ protocols::{
chat::ChatCompletionRequest, chat::ChatCompletionRequest,
classify::ClassifyRequest, classify::ClassifyRequest,
@@ -168,7 +168,7 @@ impl Router {
None => self.policy_registry.get_default_policy(), 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) // Record worker selection metric (Layer 3)
Metrics::record_worker_selection( Metrics::record_worker_selection(
@@ -2,7 +2,7 @@ use std::{collections::HashMap, sync::Arc};
use sgl_model_gateway::{ use sgl_model_gateway::{
core::{BasicWorkerBuilder, Worker, WorkerType}, core::{BasicWorkerBuilder, Worker, WorkerType},
policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy}, policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy, SelectWorkerInfo},
}; };
#[test] #[test]
@@ -40,7 +40,12 @@ fn test_backward_compatibility_with_empty_model_id() {
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker1.clone()), Arc::new(worker2.clone())]; let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker1.clone()), Arc::new(worker2.clone())];
// Select worker - should work without errors // 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"); assert!(selected.is_some(), "Should select a worker");
// Remove workers - should work without errors // Remove workers - should work without errors
@@ -97,12 +102,15 @@ fn test_mixed_model_ids() {
let default_workers: Vec<Arc<dyn Worker>> = let default_workers: Vec<Arc<dyn Worker>> =
vec![Arc::new(worker1.clone()), Arc::new(worker3.clone())]; 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"); assert!(selected.is_some(), "Should select from default workers");
let llama_workers: Vec<Arc<dyn Worker>> = let llama_workers: Vec<Arc<dyn Worker>> =
vec![Arc::new(worker2.clone()), Arc::new(worker4.clone())]; 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"); assert!(selected.is_some(), "Should select from llama-3 workers");
let all_workers: Vec<Arc<dyn Worker>> = vec![ let all_workers: Vec<Arc<dyn Worker>> = vec![
@@ -111,7 +119,7 @@ fn test_mixed_model_ids() {
Arc::new(worker3.clone()), Arc::new(worker3.clone()),
Arc::new(worker4.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"); 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"); policy.remove_worker_by_url("http://worker1:8080");
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker2.clone())]; let workers: Vec<Arc<dyn Worker>> = 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"); assert_eq!(selected, Some(0), "Should only have worker2 left");
} }