Change routing policy API to be async to support more policies (#17048)
This commit is contained in:
@@ -1,10 +1,11 @@
|
||||
use std::{sync::Arc, thread};
|
||||
use std::sync::Arc;
|
||||
|
||||
use criterion::{black_box, criterion_group, criterion_main, BenchmarkId, Criterion, Throughput};
|
||||
use smg::{
|
||||
core::{BasicWorkerBuilder, Worker, WorkerType},
|
||||
policies::{LoadBalancingPolicy, ManualPolicy, SelectWorkerInfo},
|
||||
};
|
||||
use tokio::runtime::Runtime;
|
||||
|
||||
// ============================================================================
|
||||
// Test Helpers
|
||||
@@ -22,19 +23,24 @@ fn create_workers(count: usize) -> Vec<Arc<dyn Worker>> {
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn select_with_key(policy: &ManualPolicy, workers: &[Arc<dyn Worker>], key: &str) -> Option<usize> {
|
||||
fn select_with_key(
|
||||
rt: &Runtime,
|
||||
policy: &ManualPolicy,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
key: &str,
|
||||
) -> Option<usize> {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("x-smg-routing-key", key.parse().unwrap());
|
||||
let info = SelectWorkerInfo {
|
||||
headers: Some(&headers),
|
||||
..Default::default()
|
||||
};
|
||||
policy.select_worker(workers, &info)
|
||||
rt.block_on(policy.select_worker(workers, &info))
|
||||
}
|
||||
|
||||
fn warmup_keys(policy: &ManualPolicy, workers: &[Arc<dyn Worker>], keys: &[String]) {
|
||||
fn warmup_keys(rt: &Runtime, policy: &ManualPolicy, workers: &[Arc<dyn Worker>], keys: &[String]) {
|
||||
for key in keys {
|
||||
select_with_key(policy, workers, key);
|
||||
select_with_key(rt, policy, workers, key);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -47,13 +53,14 @@ fn gen_keys(count: usize, prefix: &str) -> Vec<String> {
|
||||
// ============================================================================
|
||||
|
||||
fn bench_fast_path_hit(c: &mut Criterion) {
|
||||
let rt = Runtime::new().unwrap();
|
||||
let mut group = c.benchmark_group("manual_policy/fast_path");
|
||||
|
||||
for worker_count in [4, 16, 64, 256] {
|
||||
let policy = ManualPolicy::new();
|
||||
let workers = create_workers(worker_count);
|
||||
let keys = gen_keys(1000, "user-");
|
||||
warmup_keys(&policy, &workers, &keys);
|
||||
warmup_keys(&rt, &policy, &workers, &keys);
|
||||
|
||||
group.throughput(Throughput::Elements(1));
|
||||
group.bench_with_input(
|
||||
@@ -62,7 +69,7 @@ fn bench_fast_path_hit(c: &mut Criterion) {
|
||||
|b, _| {
|
||||
let mut idx = 0;
|
||||
b.iter(|| {
|
||||
let result = select_with_key(&policy, &workers, &keys[idx % keys.len()]);
|
||||
let result = select_with_key(&rt, &policy, &workers, &keys[idx % keys.len()]);
|
||||
idx += 1;
|
||||
black_box(result)
|
||||
});
|
||||
@@ -73,6 +80,7 @@ fn bench_fast_path_hit(c: &mut Criterion) {
|
||||
}
|
||||
|
||||
fn bench_slow_path_vacant(c: &mut Criterion) {
|
||||
let rt = Runtime::new().unwrap();
|
||||
let mut group = c.benchmark_group("manual_policy/slow_path_vacant");
|
||||
|
||||
for worker_count in [4, 16, 64, 256] {
|
||||
@@ -87,7 +95,7 @@ fn bench_slow_path_vacant(c: &mut Criterion) {
|
||||
let mut idx = 0;
|
||||
b.iter(|| {
|
||||
let key = format!("new-user-{}", idx);
|
||||
let result = select_with_key(&policy, &workers, &key);
|
||||
let result = select_with_key(&rt, &policy, &workers, &key);
|
||||
idx += 1;
|
||||
black_box(result)
|
||||
});
|
||||
@@ -98,6 +106,7 @@ fn bench_slow_path_vacant(c: &mut Criterion) {
|
||||
}
|
||||
|
||||
fn bench_no_routing_key(c: &mut Criterion) {
|
||||
let rt = Runtime::new().unwrap();
|
||||
let mut group = c.benchmark_group("manual_policy/no_routing_key");
|
||||
|
||||
for worker_count in [4, 16, 64, 256] {
|
||||
@@ -110,7 +119,7 @@ fn bench_no_routing_key(c: &mut Criterion) {
|
||||
&worker_count,
|
||||
|b, _| {
|
||||
let info = SelectWorkerInfo::default();
|
||||
b.iter(|| black_box(policy.select_worker(&workers, &info)));
|
||||
b.iter(|| black_box(rt.block_on(policy.select_worker(&workers, &info))));
|
||||
},
|
||||
);
|
||||
}
|
||||
@@ -118,6 +127,7 @@ fn bench_no_routing_key(c: &mut Criterion) {
|
||||
}
|
||||
|
||||
fn bench_failover(c: &mut Criterion) {
|
||||
let rt = Runtime::new().unwrap();
|
||||
let mut group = c.benchmark_group("manual_policy/failover");
|
||||
group.sample_size(50);
|
||||
|
||||
@@ -130,12 +140,12 @@ fn bench_failover(c: &mut Criterion) {
|
||||
|| {
|
||||
let policy = ManualPolicy::new();
|
||||
let workers = create_workers(count);
|
||||
let idx = select_with_key(&policy, &workers, "failover-test").unwrap();
|
||||
let idx = select_with_key(&rt, &policy, &workers, "failover-test").unwrap();
|
||||
workers[idx].set_healthy(false);
|
||||
(policy, workers)
|
||||
},
|
||||
|(policy, workers)| {
|
||||
black_box(select_with_key(&policy, &workers, "failover-test"))
|
||||
black_box(select_with_key(&rt, &policy, &workers, "failover-test"))
|
||||
},
|
||||
);
|
||||
},
|
||||
@@ -145,6 +155,12 @@ fn bench_failover(c: &mut Criterion) {
|
||||
}
|
||||
|
||||
fn bench_concurrent(c: &mut Criterion) {
|
||||
let rt = Arc::new(
|
||||
tokio::runtime::Builder::new_multi_thread()
|
||||
.worker_threads(4)
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let mut group = c.benchmark_group("manual_policy/concurrent");
|
||||
group.sample_size(50);
|
||||
|
||||
@@ -157,32 +173,35 @@ fn bench_concurrent(c: &mut Criterion) {
|
||||
let policy = Arc::new(ManualPolicy::new());
|
||||
let workers: Arc<Vec<Arc<dyn Worker>>> = Arc::new(create_workers(16));
|
||||
|
||||
let handles: Vec<_> = (0..threads)
|
||||
.map(|t| {
|
||||
let policy = Arc::clone(&policy);
|
||||
let workers = Arc::clone(&workers);
|
||||
thread::spawn(move || {
|
||||
for i in 0..500 {
|
||||
let key = if i % 5 == 0 {
|
||||
format!("thread{}_user{}", t, i)
|
||||
} else {
|
||||
format!("shared_user{}", i % 50)
|
||||
};
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("x-smg-routing-key", key.parse().unwrap());
|
||||
let info = SelectWorkerInfo {
|
||||
headers: Some(&headers),
|
||||
..Default::default()
|
||||
};
|
||||
let _ = black_box(policy.select_worker(&workers, &info));
|
||||
}
|
||||
rt.block_on(async {
|
||||
let handles: Vec<_> = (0..threads)
|
||||
.map(|t| {
|
||||
let policy = Arc::clone(&policy);
|
||||
let workers = Arc::clone(&workers);
|
||||
tokio::spawn(async move {
|
||||
for i in 0..500 {
|
||||
let key = if i % 5 == 0 {
|
||||
format!("thread{}_user{}", t, i)
|
||||
} else {
|
||||
format!("shared_user{}", i % 50)
|
||||
};
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("x-smg-routing-key", key.parse().unwrap());
|
||||
let info = SelectWorkerInfo {
|
||||
headers: Some(&headers),
|
||||
..Default::default()
|
||||
};
|
||||
let _ =
|
||||
black_box(policy.select_worker(&workers, &info).await);
|
||||
}
|
||||
})
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
.collect();
|
||||
|
||||
for h in handles {
|
||||
h.join().unwrap();
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
});
|
||||
});
|
||||
},
|
||||
);
|
||||
@@ -191,19 +210,20 @@ fn bench_concurrent(c: &mut Criterion) {
|
||||
}
|
||||
|
||||
fn bench_cache_size_impact(c: &mut Criterion) {
|
||||
let rt = Runtime::new().unwrap();
|
||||
let mut group = c.benchmark_group("manual_policy/cache_size");
|
||||
|
||||
for cache_size in [100, 1000, 10000, 100000] {
|
||||
let policy = ManualPolicy::new();
|
||||
let workers = create_workers(16);
|
||||
let keys = gen_keys(cache_size, "user-");
|
||||
warmup_keys(&policy, &workers, &keys);
|
||||
warmup_keys(&rt, &policy, &workers, &keys);
|
||||
|
||||
group.throughput(Throughput::Elements(1));
|
||||
group.bench_with_input(BenchmarkId::new("keys", cache_size), &cache_size, |b, _| {
|
||||
let mut idx = 0;
|
||||
b.iter(|| {
|
||||
let result = select_with_key(&policy, &workers, &keys[idx % keys.len()]);
|
||||
let result = select_with_key(&rt, &policy, &workers, &keys[idx % keys.len()]);
|
||||
idx += 1;
|
||||
black_box(result)
|
||||
});
|
||||
|
||||
@@ -5,6 +5,7 @@ use std::{
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use dashmap::DashMap;
|
||||
use rand::Rng;
|
||||
use tracing::{debug, error, info, warn};
|
||||
@@ -203,8 +204,13 @@ impl BucketPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for BucketPolicy {
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
|
||||
if healthy_indices.is_empty() {
|
||||
@@ -628,6 +634,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
// Two requests len =34 ->load balancing
|
||||
policy
|
||||
@@ -638,6 +645,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -647,6 +655,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(11)).await;
|
||||
@@ -681,6 +690,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let idx_2 = policy
|
||||
.select_worker(
|
||||
@@ -690,6 +700,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let idx_3 = policy
|
||||
.select_worker(
|
||||
@@ -699,6 +710,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx_1, 0, "Should not trigger load balancing");
|
||||
assert_ne!(idx_2, idx_3, "Should trigger load balancing");
|
||||
@@ -723,6 +735,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap(); // worker1: 20
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -732,6 +745,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap(); // worker1: 8
|
||||
|
||||
// Next request should not use bucket scheduling (no load balancing)
|
||||
@@ -743,6 +757,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
idx, 0,
|
||||
@@ -768,6 +783,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap(); // worker1: 15
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -777,6 +793,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap(); // worker2: 5
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -786,6 +803,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap(); // worker3: 10
|
||||
|
||||
// Next request should use bucket scheduling (load balancing)
|
||||
@@ -797,6 +815,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
idx, 0,
|
||||
@@ -870,6 +889,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -879,6 +899,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -888,6 +909,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -897,6 +919,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -906,6 +929,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -915,6 +939,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
@@ -951,6 +976,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -960,6 +986,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -969,6 +996,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -978,6 +1006,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -987,6 +1016,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -996,6 +1026,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
@@ -1087,6 +1118,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
@@ -1120,6 +1152,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
@@ -1209,6 +1242,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1218,6 +1252,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1227,6 +1262,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1236,6 +1272,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1245,6 +1282,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1254,6 +1292,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
@@ -1287,6 +1326,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1296,6 +1336,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1305,6 +1346,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1314,6 +1356,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1323,6 +1366,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
policy
|
||||
.select_worker(
|
||||
@@ -1332,6 +1376,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
tokio::time::sleep(Duration::from_secs(4)).await;
|
||||
|
||||
@@ -61,6 +61,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use dashmap::DashMap;
|
||||
use rand::Rng;
|
||||
use tracing::{debug, warn};
|
||||
@@ -347,8 +348,13 @@ impl CacheAwarePolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for CacheAwarePolicy {
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let request_text = info.request_text;
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
|
||||
@@ -508,8 +514,8 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
|
||||
#[test]
|
||||
fn test_cache_aware_with_balanced_load() {
|
||||
#[tokio::test]
|
||||
async fn test_cache_aware_with_balanced_load() {
|
||||
// Create policy without eviction thread for testing
|
||||
let config = CacheAwareConfig {
|
||||
eviction_interval_secs: 0, // Disable eviction thread
|
||||
@@ -543,6 +549,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Same request should go to same worker (cache hit)
|
||||
@@ -554,6 +561,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx1, idx2);
|
||||
|
||||
@@ -566,12 +574,13 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx1, idx3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cache_aware_with_imbalanced_load() {
|
||||
#[tokio::test]
|
||||
async fn test_cache_aware_with_imbalanced_load() {
|
||||
let policy = CacheAwarePolicy::with_config(CacheAwareConfig {
|
||||
cache_threshold: 0.5,
|
||||
balance_abs_threshold: 5,
|
||||
@@ -602,13 +611,13 @@ mod tests {
|
||||
..Default::default()
|
||||
};
|
||||
for _ in 0..5 {
|
||||
let idx = policy.select_worker(&workers, &info).unwrap();
|
||||
let idx = policy.select_worker(&workers, &info).await.unwrap();
|
||||
assert_eq!(idx, 1); // Should always pick worker2
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cache_aware_worker_removal() {
|
||||
#[tokio::test]
|
||||
async fn test_cache_aware_worker_removal() {
|
||||
let config = CacheAwareConfig {
|
||||
eviction_interval_secs: 0, // Disable eviction thread
|
||||
..Default::default()
|
||||
@@ -630,20 +639,24 @@ mod tests {
|
||||
policy.init_workers(&workers);
|
||||
|
||||
// Route some requests
|
||||
policy.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test1"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
policy.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test2"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
policy
|
||||
.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test1"),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await;
|
||||
policy
|
||||
.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test2"),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
// Remove a worker
|
||||
policy.remove_worker_by_url("http://w1:8000");
|
||||
@@ -658,12 +671,13 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cache_aware_sync_tree_operation_to_mesh() {
|
||||
#[tokio::test]
|
||||
async fn test_cache_aware_sync_tree_operation_to_mesh() {
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::mesh::{stores::StateStores, sync::MeshSyncManager};
|
||||
@@ -696,6 +710,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Verify tree operation was synced to mesh (under UNKNOWN_MODEL_ID since no model was specified)
|
||||
@@ -841,8 +856,8 @@ mod tests {
|
||||
let _ = tree_state;
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cache_aware_without_mesh() {
|
||||
#[tokio::test]
|
||||
async fn test_cache_aware_without_mesh() {
|
||||
let config = CacheAwareConfig {
|
||||
eviction_interval_secs: 0,
|
||||
..Default::default()
|
||||
@@ -867,6 +882,7 @@ mod tests {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx, 0);
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rand::Rng as _;
|
||||
|
||||
use super::{LoadBalancingPolicy, SelectWorkerInfo};
|
||||
@@ -167,8 +168,13 @@ impl ConsistentHashingPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for ConsistentHashingPolicy {
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let (result, branch) = self.select_worker_impl(workers, info);
|
||||
Metrics::record_worker_consistent_hashing_policy_branch(branch.as_str());
|
||||
result
|
||||
@@ -214,8 +220,8 @@ mod tests {
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_consistent_routing() {
|
||||
#[tokio::test]
|
||||
async fn test_consistent_routing() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
|
||||
|
||||
@@ -236,8 +242,8 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_different_keys_distribute() {
|
||||
#[tokio::test]
|
||||
async fn test_different_keys_distribute() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
|
||||
|
||||
@@ -255,8 +261,8 @@ mod tests {
|
||||
assert!(distribution.len() > 1, "Should distribute across workers");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_target_worker_hit() {
|
||||
#[tokio::test]
|
||||
async fn test_target_worker_hit() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
|
||||
|
||||
@@ -271,8 +277,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::TargetWorkerHit);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_target_worker_miss_out_of_bounds() {
|
||||
#[tokio::test]
|
||||
async fn test_target_worker_miss_out_of_bounds() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
|
||||
|
||||
@@ -287,8 +293,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::TargetWorkerMiss);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_target_worker_miss_unhealthy() {
|
||||
#[tokio::test]
|
||||
async fn test_target_worker_miss_unhealthy() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
|
||||
workers[1].set_healthy(false);
|
||||
@@ -304,8 +310,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::TargetWorkerMiss);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_target_worker_priority_over_routing_key() {
|
||||
#[tokio::test]
|
||||
async fn test_target_worker_priority_over_routing_key() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
|
||||
|
||||
@@ -323,8 +329,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::TargetWorkerHit);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_fallback_random_distribution() {
|
||||
#[tokio::test]
|
||||
async fn test_fallback_random_distribution() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
|
||||
|
||||
@@ -345,8 +351,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_no_healthy_workers() {
|
||||
#[tokio::test]
|
||||
async fn test_no_healthy_workers() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000"]);
|
||||
workers[0].set_healthy(false);
|
||||
@@ -362,8 +368,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::NoHealthyWorkers);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_workers() {
|
||||
#[tokio::test]
|
||||
async fn test_empty_workers() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![];
|
||||
|
||||
@@ -373,8 +379,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::NoHealthyWorkers);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_consistent_hash_minimal_redistribution() {
|
||||
#[tokio::test]
|
||||
async fn test_consistent_hash_minimal_redistribution() {
|
||||
// Test that consistent hashing moves fewer keys than random redistribution
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&[
|
||||
@@ -439,8 +445,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routing_key_failover_and_recovery() {
|
||||
#[tokio::test]
|
||||
async fn test_routing_key_failover_and_recovery() {
|
||||
// Test that when a worker fails, keys move to another worker,
|
||||
// and when it recovers, keys return to the original worker
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
@@ -491,8 +497,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_routing_key_uses_fallback() {
|
||||
#[tokio::test]
|
||||
async fn test_empty_routing_key_uses_fallback() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
let workers = create_workers(&["http://w1:8000", "http://w2:8000"]);
|
||||
|
||||
@@ -507,8 +513,8 @@ mod tests {
|
||||
assert_eq!(branch, Branch::RandomFallback);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_policy_name() {
|
||||
#[tokio::test]
|
||||
async fn test_policy_name() {
|
||||
let policy = ConsistentHashingPolicy::new();
|
||||
assert_eq!(policy.name(), "consistent_hashing");
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
|
||||
use std::{sync::Arc, time::Instant};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use dashmap::{mapref::entry::Entry, DashMap};
|
||||
use rand::Rng;
|
||||
use tracing::info;
|
||||
@@ -194,7 +195,7 @@ impl ManualPolicy {
|
||||
fn select_worker_impl(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo,
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> (Option<usize>, ExecutionBranch) {
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
if healthy_indices.is_empty() {
|
||||
@@ -213,8 +214,13 @@ impl ManualPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for ManualPolicy {
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let (result, branch) = self.select_worker_impl(workers, info);
|
||||
Metrics::record_worker_manual_policy_branch(branch.as_str());
|
||||
Metrics::set_manual_policy_cache_entries(self.routing_map.len());
|
||||
|
||||
@@ -5,6 +5,8 @@
|
||||
|
||||
use std::{fmt::Debug, sync::Arc};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::{
|
||||
core::{HashRing, Worker},
|
||||
mesh::OptionalMeshSyncManager,
|
||||
@@ -38,6 +40,7 @@ pub use tree::PrefixMatchResult;
|
||||
///
|
||||
/// This trait provides a unified interface for implementing routing algorithms
|
||||
/// that can work with both regular single-worker selection and PD dual-worker selection.
|
||||
#[async_trait]
|
||||
pub trait LoadBalancingPolicy: Send + Sync + Debug {
|
||||
/// Select a single worker from the available workers
|
||||
///
|
||||
@@ -47,7 +50,11 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug {
|
||||
/// # Arguments
|
||||
/// * `workers` - Available workers to select from
|
||||
/// * `info` - Additional information for routing decisions
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize>;
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize>;
|
||||
|
||||
/// Update policy state after request completion
|
||||
///
|
||||
@@ -173,8 +180,8 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
|
||||
#[test]
|
||||
fn test_get_healthy_worker_indices() {
|
||||
#[tokio::test]
|
||||
async fn test_get_healthy_worker_indices() {
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
BasicWorkerBuilder::new("http://w1:8000")
|
||||
|
||||
@@ -5,6 +5,7 @@ use std::{
|
||||
sync::{Arc, RwLock},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rand::Rng;
|
||||
use tracing::debug;
|
||||
|
||||
@@ -29,11 +30,12 @@ impl PowerOfTwoPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for PowerOfTwoPolicy {
|
||||
fn select_worker(
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
_info: &SelectWorkerInfo,
|
||||
_info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
|
||||
@@ -130,8 +132,8 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
|
||||
#[test]
|
||||
fn test_power_of_two_selection() {
|
||||
#[tokio::test]
|
||||
async fn test_power_of_two_selection() {
|
||||
let policy = PowerOfTwoPolicy::new();
|
||||
let worker1 = BasicWorkerBuilder::new("http://w1:8000")
|
||||
.worker_type(WorkerType::Regular)
|
||||
@@ -159,7 +161,7 @@ mod tests {
|
||||
let mut selected_counts = [0; 3];
|
||||
let info = SelectWorkerInfo::default();
|
||||
for _ in 0..100 {
|
||||
if let Some(idx) = policy.select_worker(&workers, &info) {
|
||||
if let Some(idx) = policy.select_worker(&workers, &info).await {
|
||||
selected_counts[idx] += 1;
|
||||
}
|
||||
}
|
||||
@@ -169,8 +171,8 @@ mod tests {
|
||||
assert!(selected_counts[1] > selected_counts[0]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_power_of_two_with_cached_loads() {
|
||||
#[tokio::test]
|
||||
async fn test_power_of_two_with_cached_loads() {
|
||||
let policy = PowerOfTwoPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -195,7 +197,7 @@ mod tests {
|
||||
let mut w2_selected = 0;
|
||||
let info = SelectWorkerInfo::default();
|
||||
for _ in 0..50 {
|
||||
if let Some(idx) = policy.select_worker(&workers, &info) {
|
||||
if let Some(idx) = policy.select_worker(&workers, &info).await {
|
||||
if idx == 1 {
|
||||
w2_selected += 1;
|
||||
}
|
||||
@@ -206,8 +208,8 @@ mod tests {
|
||||
assert!(w2_selected > 35); // Should win most of the time
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_power_of_two_single_worker() {
|
||||
#[tokio::test]
|
||||
async fn test_power_of_two_single_worker() {
|
||||
let policy = PowerOfTwoPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(
|
||||
BasicWorkerBuilder::new("http://w1:8000")
|
||||
@@ -217,13 +219,15 @@ mod tests {
|
||||
|
||||
// With single worker, should always select it
|
||||
assert_eq!(
|
||||
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||
policy
|
||||
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||
.await,
|
||||
Some(0)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_reproduce_incompatible_metric_bug() {
|
||||
#[tokio::test]
|
||||
async fn test_reproduce_incompatible_metric_bug() {
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
@@ -257,6 +261,7 @@ mod tests {
|
||||
// 5. Run selection
|
||||
let selected_idx = policy
|
||||
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||
.await
|
||||
.expect("Should select a worker");
|
||||
|
||||
// 6. Verify the Fix
|
||||
@@ -280,8 +285,8 @@ mod tests {
|
||||
"The policy failed to handle incompatible metrics. Should select idle Worker A."
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn test_power_of_two_edge_cases() {
|
||||
#[tokio::test]
|
||||
async fn test_power_of_two_edge_cases() {
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
@@ -314,6 +319,7 @@ mod tests {
|
||||
|
||||
let idx_1 = policy
|
||||
.select_worker(&workers_1, &SelectWorkerInfo::default())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
idx_1, 0,
|
||||
@@ -335,6 +341,7 @@ mod tests {
|
||||
|
||||
let idx_2 = policy
|
||||
.select_worker(&workers_2, &SelectWorkerInfo::default())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx_2, 1, "Partial Fail 1 Failed: Should fallback to requests and select Worker B (fewer requests)");
|
||||
|
||||
@@ -353,6 +360,7 @@ mod tests {
|
||||
|
||||
let idx_3 = policy
|
||||
.select_worker(&workers_3, &SelectWorkerInfo::default())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(idx_3, 0, "Partial Fail 2 Failed: Should fallback to requests and select Worker A (fewer requests)");
|
||||
|
||||
@@ -369,6 +377,7 @@ mod tests {
|
||||
|
||||
let idx_4 = policy
|
||||
.select_worker(&workers_4, &SelectWorkerInfo::default())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
idx_4, 1,
|
||||
|
||||
@@ -221,8 +221,13 @@ impl PrefixHashPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl LoadBalancingPolicy for PrefixHashPolicy {
|
||||
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let (result, branch) = self.select_worker_impl(workers, info);
|
||||
Metrics::record_worker_prefix_hash_policy_branch(branch.as_str());
|
||||
result
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rand::Rng;
|
||||
|
||||
use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
|
||||
@@ -19,11 +20,12 @@ impl RandomPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for RandomPolicy {
|
||||
fn select_worker(
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
_info: &SelectWorkerInfo,
|
||||
_info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
|
||||
@@ -53,8 +55,8 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
|
||||
#[test]
|
||||
fn test_random_selection() {
|
||||
#[tokio::test]
|
||||
async fn test_random_selection() {
|
||||
let policy = RandomPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -76,18 +78,20 @@ mod tests {
|
||||
|
||||
let mut counts = HashMap::new();
|
||||
for _ in 0..100 {
|
||||
if let Some(idx) = policy.select_worker(&workers, &SelectWorkerInfo::default()) {
|
||||
if let Some(idx) = policy
|
||||
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||
.await
|
||||
{
|
||||
*counts.entry(idx).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
|
||||
// All workers should be selected at least once
|
||||
assert_eq!(counts.len(), 3);
|
||||
assert!(counts.values().all(|&count| count > 0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_random_with_unhealthy_workers() {
|
||||
#[tokio::test]
|
||||
async fn test_random_with_unhealthy_workers() {
|
||||
let policy = RandomPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -102,20 +106,20 @@ mod tests {
|
||||
),
|
||||
];
|
||||
|
||||
// Mark first worker as unhealthy
|
||||
workers[0].set_healthy(false);
|
||||
|
||||
// Should always select the healthy worker (index 1)
|
||||
for _ in 0..10 {
|
||||
assert_eq!(
|
||||
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||
policy
|
||||
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||
.await,
|
||||
Some(1)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_random_no_healthy_workers() {
|
||||
#[tokio::test]
|
||||
async fn test_random_no_healthy_workers() {
|
||||
let policy = RandomPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(
|
||||
BasicWorkerBuilder::new("http://w1:8000")
|
||||
@@ -125,7 +129,9 @@ mod tests {
|
||||
|
||||
workers[0].set_healthy(false);
|
||||
assert_eq!(
|
||||
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||
policy
|
||||
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||
.await,
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
@@ -449,8 +449,8 @@ impl std::fmt::Debug for PolicyRegistry {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_policy_registry_basic() {
|
||||
#[tokio::test]
|
||||
async fn test_policy_registry_basic() {
|
||||
let registry = PolicyRegistry::new(PolicyConfig::RoundRobin);
|
||||
|
||||
// First worker of a model sets the policy
|
||||
@@ -476,8 +476,8 @@ mod tests {
|
||||
assert_eq!(*counts.get("gpt-4").unwrap(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_policy_registry_cleanup() {
|
||||
#[tokio::test]
|
||||
async fn test_policy_registry_cleanup() {
|
||||
let registry = PolicyRegistry::new(PolicyConfig::RoundRobin);
|
||||
|
||||
// Add workers
|
||||
@@ -496,8 +496,8 @@ mod tests {
|
||||
assert_eq!(registry.get_worker_counts().get("llama-3"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_policy() {
|
||||
#[tokio::test]
|
||||
async fn test_default_policy() {
|
||||
let registry = PolicyRegistry::new(PolicyConfig::RoundRobin);
|
||||
|
||||
// No hint, no template - uses default
|
||||
|
||||
@@ -5,6 +5,8 @@ use std::sync::{
|
||||
Arc,
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
|
||||
use crate::core::Worker;
|
||||
|
||||
@@ -24,11 +26,12 @@ impl RoundRobinPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LoadBalancingPolicy for RoundRobinPolicy {
|
||||
fn select_worker(
|
||||
async fn select_worker(
|
||||
&self,
|
||||
workers: &[Arc<dyn Worker>],
|
||||
_info: &SelectWorkerInfo,
|
||||
_info: &SelectWorkerInfo<'_>,
|
||||
) -> Option<usize> {
|
||||
let healthy_indices = get_healthy_worker_indices(workers);
|
||||
|
||||
@@ -61,8 +64,8 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
|
||||
#[test]
|
||||
fn test_round_robin_selection() {
|
||||
#[tokio::test]
|
||||
async fn test_round_robin_selection() {
|
||||
let policy = RoundRobinPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -82,17 +85,16 @@ mod tests {
|
||||
),
|
||||
];
|
||||
|
||||
// Should select workers in order: 0, 1, 2, 0, 1, 2, ...
|
||||
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));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(1));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(2));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_round_robin_with_unhealthy_workers() {
|
||||
#[tokio::test]
|
||||
async fn test_round_robin_with_unhealthy_workers() {
|
||||
let policy = RoundRobinPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -112,19 +114,17 @@ mod tests {
|
||||
),
|
||||
];
|
||||
|
||||
// Mark middle worker as unhealthy
|
||||
workers[1].set_healthy(false);
|
||||
|
||||
// Should skip unhealthy worker: 0, 2, 0, 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));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(2));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_round_robin_reset() {
|
||||
#[tokio::test]
|
||||
async fn test_round_robin_reset() {
|
||||
let policy = RoundRobinPolicy::new();
|
||||
let workers: Vec<Arc<dyn Worker>> = vec![
|
||||
Arc::new(
|
||||
@@ -139,13 +139,11 @@ mod tests {
|
||||
),
|
||||
];
|
||||
|
||||
// Advance the counter
|
||||
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).await, Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(1));
|
||||
|
||||
// Reset should start from beginning
|
||||
policy.reset();
|
||||
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||
assert_eq!(policy.select_worker(&workers, &info).await, Some(0));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,12 +78,10 @@ impl PipelineStage for WorkerSelectionStage {
|
||||
|
||||
let workers = match self.mode {
|
||||
WorkerSelectionMode::Regular => {
|
||||
match self.select_single_worker(
|
||||
ctx.input.model_id.as_deref(),
|
||||
text,
|
||||
tokens,
|
||||
headers,
|
||||
) {
|
||||
match self
|
||||
.select_single_worker(ctx.input.model_id.as_deref(), text, tokens, headers)
|
||||
.await
|
||||
{
|
||||
Some(w) => WorkerSelection::Single { worker: w },
|
||||
None => {
|
||||
let model = ctx.input.model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID);
|
||||
@@ -101,7 +99,10 @@ impl PipelineStage for WorkerSelectionStage {
|
||||
}
|
||||
}
|
||||
WorkerSelectionMode::PrefillDecode => {
|
||||
match self.select_pd_pair(ctx.input.model_id.as_deref(), text, tokens, headers) {
|
||||
match self
|
||||
.select_pd_pair(ctx.input.model_id.as_deref(), text, tokens, headers)
|
||||
.await
|
||||
{
|
||||
Some((prefill, decode)) => WorkerSelection::Dual { prefill, decode },
|
||||
None => {
|
||||
let model = ctx.input.model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID);
|
||||
@@ -130,7 +131,7 @@ impl PipelineStage for WorkerSelectionStage {
|
||||
}
|
||||
|
||||
impl WorkerSelectionStage {
|
||||
fn select_single_worker(
|
||||
async fn select_single_worker(
|
||||
&self,
|
||||
model_id: Option<&str>,
|
||||
text: Option<&str>,
|
||||
@@ -166,15 +167,17 @@ impl WorkerSelectionStage {
|
||||
.get_hash_ring(model_id.unwrap_or(UNKNOWN_MODEL_ID));
|
||||
|
||||
// Select worker using the policy
|
||||
let idx = policy.select_worker(
|
||||
&available,
|
||||
&SelectWorkerInfo {
|
||||
request_text: text,
|
||||
tokens,
|
||||
headers,
|
||||
hash_ring,
|
||||
},
|
||||
)?;
|
||||
let idx = policy
|
||||
.select_worker(
|
||||
&available,
|
||||
&SelectWorkerInfo {
|
||||
request_text: text,
|
||||
tokens,
|
||||
headers,
|
||||
hash_ring,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
let selected = available[idx].clone();
|
||||
|
||||
// Record worker selection metric
|
||||
@@ -188,7 +191,7 @@ impl WorkerSelectionStage {
|
||||
Some(selected)
|
||||
}
|
||||
|
||||
fn select_pd_pair(
|
||||
async fn select_pd_pair(
|
||||
&self,
|
||||
model_id: Option<&str>,
|
||||
text: Option<&str>,
|
||||
@@ -244,8 +247,8 @@ impl WorkerSelectionStage {
|
||||
headers,
|
||||
hash_ring,
|
||||
};
|
||||
let prefill_idx = policy.select_worker(&available_prefill, &info)?;
|
||||
let decode_idx = policy.select_worker(&available_decode, &info)?;
|
||||
let prefill_idx = policy.select_worker(&available_prefill, &info).await?;
|
||||
let decode_idx = policy.select_worker(&available_decode, &info).await?;
|
||||
|
||||
let model = model_id.unwrap_or(UNKNOWN_MODEL_ID);
|
||||
let policy_name = policy.name();
|
||||
|
||||
@@ -747,7 +747,8 @@ impl PDRouter {
|
||||
headers,
|
||||
hash_ring.clone(),
|
||||
"prefill",
|
||||
)?;
|
||||
)
|
||||
.await?;
|
||||
|
||||
let decode = Self::pick_worker_by_policy_arc(
|
||||
&decode_workers,
|
||||
@@ -756,7 +757,8 @@ impl PDRouter {
|
||||
headers,
|
||||
hash_ring,
|
||||
"decode",
|
||||
)?;
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Record worker selection metrics (Layer 3)
|
||||
let model = model_id.unwrap_or(UNKNOWN_MODEL_ID);
|
||||
@@ -776,7 +778,7 @@ impl PDRouter {
|
||||
Ok((prefill, decode))
|
||||
}
|
||||
|
||||
fn pick_worker_by_policy_arc(
|
||||
async fn pick_worker_by_policy_arc(
|
||||
workers: &[Arc<dyn Worker>],
|
||||
policy: &dyn LoadBalancingPolicy,
|
||||
request_text: Option<&str>,
|
||||
@@ -814,6 +816,7 @@ impl PDRouter {
|
||||
hash_ring,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.ok_or_else(|| {
|
||||
format!(
|
||||
"Policy {} failed to select a {} worker",
|
||||
|
||||
@@ -130,7 +130,7 @@ impl Router {
|
||||
}
|
||||
|
||||
/// Select worker for a specific model considering circuit breaker state
|
||||
fn select_worker_for_model(
|
||||
async fn select_worker_for_model(
|
||||
&self,
|
||||
model_id: Option<&str>,
|
||||
text: Option<&str>,
|
||||
@@ -167,15 +167,17 @@ impl Router {
|
||||
.worker_registry
|
||||
.get_hash_ring(effective_model_id.unwrap_or(UNKNOWN_MODEL_ID));
|
||||
|
||||
let idx = policy.select_worker(
|
||||
&available,
|
||||
&SelectWorkerInfo {
|
||||
request_text: text,
|
||||
tokens: None, // HTTP doesn't have tokens, use gRPC for PrefixHash
|
||||
headers,
|
||||
hash_ring,
|
||||
},
|
||||
)?;
|
||||
let idx = policy
|
||||
.select_worker(
|
||||
&available,
|
||||
&SelectWorkerInfo {
|
||||
request_text: text,
|
||||
tokens: None, // HTTP doesn't have tokens, use gRPC for PrefixHash
|
||||
headers,
|
||||
hash_ring,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Record worker selection metric (Layer 3)
|
||||
Metrics::record_worker_selection(
|
||||
@@ -276,7 +278,10 @@ impl Router {
|
||||
is_stream: bool,
|
||||
text: &str,
|
||||
) -> Response {
|
||||
let worker = match self.select_worker_for_model(model_id, Some(text), headers) {
|
||||
let worker = match self
|
||||
.select_worker_for_model(model_id, Some(text), headers)
|
||||
.await
|
||||
{
|
||||
Some(w) => w,
|
||||
None => {
|
||||
return error::service_unavailable(
|
||||
|
||||
@@ -372,13 +372,13 @@ async fn test_rate_limit_window_reset() {
|
||||
manager.sync_rate_limit_inc(GLOBAL_RATE_LIMIT_COUNTER_KEY.to_string(), 50);
|
||||
let value_before = manager.get_rate_limit_value(GLOBAL_RATE_LIMIT_COUNTER_KEY);
|
||||
// Value may be None if not owner, or Some if owner
|
||||
if value_before.is_some() {
|
||||
assert!(value_before.unwrap() > 0);
|
||||
if let Some(val) = value_before {
|
||||
assert!(val > 0);
|
||||
|
||||
// Reset counter
|
||||
manager.reset_global_rate_limit_counter();
|
||||
let value_after = manager.get_rate_limit_value(GLOBAL_RATE_LIMIT_COUNTER_KEY);
|
||||
// Should be reset
|
||||
assert!(value_after.is_none() || value_after.unwrap() <= 0);
|
||||
assert!(value_after.is_none() || value_after.unwrap_or(0) <= 0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,8 +5,8 @@ use smg::{
|
||||
policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy, SelectWorkerInfo},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn test_backward_compatibility_with_empty_model_id() {
|
||||
#[tokio::test]
|
||||
async fn test_backward_compatibility_with_empty_model_id() {
|
||||
let config = CacheAwareConfig {
|
||||
cache_threshold: 0.5,
|
||||
balance_abs_threshold: 2,
|
||||
@@ -40,13 +40,15 @@ 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,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let selected = policy
|
||||
.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await;
|
||||
assert!(selected.is_some(), "Should select a worker");
|
||||
|
||||
// Remove workers - should work without errors
|
||||
@@ -54,8 +56,8 @@ fn test_backward_compatibility_with_empty_model_id() {
|
||||
policy.remove_worker(&worker2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mixed_model_ids() {
|
||||
#[tokio::test]
|
||||
async fn test_mixed_model_ids() {
|
||||
let config = CacheAwareConfig {
|
||||
cache_threshold: 0.5,
|
||||
balance_abs_threshold: 2,
|
||||
@@ -107,12 +109,12 @@ fn test_mixed_model_ids() {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
};
|
||||
let selected = policy.select_worker(&default_workers, &info);
|
||||
let selected = policy.select_worker(&default_workers, &info).await;
|
||||
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, &info);
|
||||
let selected = policy.select_worker(&llama_workers, &info).await;
|
||||
assert!(selected.is_some(), "Should select from llama-3 workers");
|
||||
|
||||
let all_workers: Vec<Arc<dyn Worker>> = vec![
|
||||
@@ -121,12 +123,12 @@ fn test_mixed_model_ids() {
|
||||
Arc::new(worker3.clone()),
|
||||
Arc::new(worker4.clone()),
|
||||
];
|
||||
let selected = policy.select_worker(&all_workers, &info);
|
||||
let selected = policy.select_worker(&all_workers, &info).await;
|
||||
assert!(selected.is_some(), "Should select from all workers");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_remove_worker_by_url_backward_compat() {
|
||||
#[tokio::test]
|
||||
async fn test_remove_worker_by_url_backward_compat() {
|
||||
let config = CacheAwareConfig::default();
|
||||
let policy = CacheAwarePolicy::with_config(config);
|
||||
|
||||
@@ -154,12 +156,14 @@ 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,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let selected = policy
|
||||
.select_worker(
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test"),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await;
|
||||
assert_eq!(selected, Some(0), "Should only have worker2 left");
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user