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
@@ -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<Arc<dyn Worker>> = 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<Arc<dyn Worker>> =
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<Arc<dyn Worker>> =
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<Arc<dyn Worker>> = 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<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");
}