Change routing policy API to be async to support more policies (#17048)

This commit is contained in:
fzyzcjy
2026-01-17 17:12:51 +08:00
committed by GitHub
parent d2c863878c
commit 9c2530642c
16 changed files with 350 additions and 217 deletions

View File

@@ -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)
});

View File

@@ -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;

View File

@@ -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);
}

View File

@@ -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");
}

View File

@@ -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());

View File

@@ -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")

View File

@@ -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,

View File

@@ -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

View File

@@ -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
);
}

View File

@@ -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

View File

@@ -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));
}
}

View File

@@ -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();

View File

@@ -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",

View File

@@ -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(

View File

@@ -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);
}
}

View File

@@ -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");
}