From 9c2530642ca4df047ec79bdec51d723c8877b88d Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sat, 17 Jan 2026 17:12:51 +0800 Subject: [PATCH] Change routing policy API to be async to support more policies (#17048) --- .../benches/manual_policy_benchmark.rs | 94 +++++++++++-------- sgl-model-gateway/src/policies/bucket.rs | 47 +++++++++- sgl-model-gateway/src/policies/cache_aware.rs | 68 +++++++++----- .../src/policies/consistent_hashing.rs | 60 ++++++------ sgl-model-gateway/src/policies/manual.rs | 10 +- sgl-model-gateway/src/policies/mod.rs | 13 ++- .../src/policies/power_of_two.rs | 39 +++++--- sgl-model-gateway/src/policies/prefix_hash.rs | 7 +- sgl-model-gateway/src/policies/random.rs | 34 ++++--- sgl-model-gateway/src/policies/registry.rs | 12 +-- sgl-model-gateway/src/policies/round_robin.rs | 48 +++++----- .../grpc/common/stages/worker_selection.rs | 43 +++++---- .../src/routers/http/pd_router.rs | 9 +- sgl-model-gateway/src/routers/http/router.rs | 27 +++--- .../tests/mesh_integration_test.rs | 6 +- .../cache_aware_backward_compat_test.rs | 50 +++++----- 16 files changed, 350 insertions(+), 217 deletions(-) diff --git a/sgl-model-gateway/benches/manual_policy_benchmark.rs b/sgl-model-gateway/benches/manual_policy_benchmark.rs index 4b98e77d0..de2701153 100644 --- a/sgl-model-gateway/benches/manual_policy_benchmark.rs +++ b/sgl-model-gateway/benches/manual_policy_benchmark.rs @@ -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> { .collect() } -fn select_with_key(policy: &ManualPolicy, workers: &[Arc], key: &str) -> Option { +fn select_with_key( + rt: &Runtime, + policy: &ManualPolicy, + workers: &[Arc], + key: &str, +) -> Option { 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], keys: &[String]) { +fn warmup_keys(rt: &Runtime, policy: &ManualPolicy, workers: &[Arc], 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 { // ============================================================================ 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>> = 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) }); diff --git a/sgl-model-gateway/src/policies/bucket.rs b/sgl-model-gateway/src/policies/bucket.rs index 3665fb361..b2e2677a5 100644 --- a/sgl-model-gateway/src/policies/bucket.rs +++ b/sgl-model-gateway/src/policies/bucket.rs @@ -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], info: &SelectWorkerInfo) -> Option { + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option { 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; diff --git a/sgl-model-gateway/src/policies/cache_aware.rs b/sgl-model-gateway/src/policies/cache_aware.rs index 010de4c1e..4fb446106 100644 --- a/sgl-model-gateway/src/policies/cache_aware.rs +++ b/sgl-model-gateway/src/policies/cache_aware.rs @@ -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], info: &SelectWorkerInfo) -> Option { + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option { 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); } diff --git a/sgl-model-gateway/src/policies/consistent_hashing.rs b/sgl-model-gateway/src/policies/consistent_hashing.rs index b9fa1c29d..2e95e8d73 100644 --- a/sgl-model-gateway/src/policies/consistent_hashing.rs +++ b/sgl-model-gateway/src/policies/consistent_hashing.rs @@ -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], info: &SelectWorkerInfo) -> Option { + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option { 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> = 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"); } diff --git a/sgl-model-gateway/src/policies/manual.rs b/sgl-model-gateway/src/policies/manual.rs index bc1f4a7bd..9aa7a1840 100644 --- a/sgl-model-gateway/src/policies/manual.rs +++ b/sgl-model-gateway/src/policies/manual.rs @@ -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], - info: &SelectWorkerInfo, + info: &SelectWorkerInfo<'_>, ) -> (Option, 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], info: &SelectWorkerInfo) -> Option { + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option { 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()); diff --git a/sgl-model-gateway/src/policies/mod.rs b/sgl-model-gateway/src/policies/mod.rs index eb9f91efd..fae218a21 100644 --- a/sgl-model-gateway/src/policies/mod.rs +++ b/sgl-model-gateway/src/policies/mod.rs @@ -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], info: &SelectWorkerInfo) -> Option; + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option; /// 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> = vec![ Arc::new( BasicWorkerBuilder::new("http://w1:8000") diff --git a/sgl-model-gateway/src/policies/power_of_two.rs b/sgl-model-gateway/src/policies/power_of_two.rs index 425e5a3d1..81cbc5c4d 100644 --- a/sgl-model-gateway/src/policies/power_of_two.rs +++ b/sgl-model-gateway/src/policies/power_of_two.rs @@ -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], - _info: &SelectWorkerInfo, + _info: &SelectWorkerInfo<'_>, ) -> Option { 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> = 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> = 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, diff --git a/sgl-model-gateway/src/policies/prefix_hash.rs b/sgl-model-gateway/src/policies/prefix_hash.rs index eeda785b0..8d79df1f2 100644 --- a/sgl-model-gateway/src/policies/prefix_hash.rs +++ b/sgl-model-gateway/src/policies/prefix_hash.rs @@ -221,8 +221,13 @@ impl PrefixHashPolicy { } } +#[async_trait::async_trait] impl LoadBalancingPolicy for PrefixHashPolicy { - fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { + async fn select_worker( + &self, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + ) -> Option { let (result, branch) = self.select_worker_impl(workers, info); Metrics::record_worker_prefix_hash_policy_branch(branch.as_str()); result diff --git a/sgl-model-gateway/src/policies/random.rs b/sgl-model-gateway/src/policies/random.rs index 0401377a7..6e8b000f2 100644 --- a/sgl-model-gateway/src/policies/random.rs +++ b/sgl-model-gateway/src/policies/random.rs @@ -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], - _info: &SelectWorkerInfo, + _info: &SelectWorkerInfo<'_>, ) -> Option { 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> = 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> = 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> = 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 ); } diff --git a/sgl-model-gateway/src/policies/registry.rs b/sgl-model-gateway/src/policies/registry.rs index 5c3d77995..19f6435b5 100644 --- a/sgl-model-gateway/src/policies/registry.rs +++ b/sgl-model-gateway/src/policies/registry.rs @@ -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 diff --git a/sgl-model-gateway/src/policies/round_robin.rs b/sgl-model-gateway/src/policies/round_robin.rs index 24dba917f..f172b8c0d 100644 --- a/sgl-model-gateway/src/policies/round_robin.rs +++ b/sgl-model-gateway/src/policies/round_robin.rs @@ -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], - _info: &SelectWorkerInfo, + _info: &SelectWorkerInfo<'_>, ) -> Option { 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> = 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> = 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> = 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)); } } diff --git a/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs b/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs index f98e6f53d..af443a9d6 100644 --- a/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs +++ b/sgl-model-gateway/src/routers/grpc/common/stages/worker_selection.rs @@ -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(); diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 0825ac110..b3a45d66b 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -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], 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", diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index e142ea7b0..0fbf2e422 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -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( diff --git a/sgl-model-gateway/tests/mesh_integration_test.rs b/sgl-model-gateway/tests/mesh_integration_test.rs index 2d7cb3bad..2aa4d5d0c 100644 --- a/sgl-model-gateway/tests/mesh_integration_test.rs +++ b/sgl-model-gateway/tests/mesh_integration_test.rs @@ -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); } } diff --git a/sgl-model-gateway/tests/routing/cache_aware_backward_compat_test.rs b/sgl-model-gateway/tests/routing/cache_aware_backward_compat_test.rs index c9f746cea..8e8365185 100644 --- a/sgl-model-gateway/tests/routing/cache_aware_backward_compat_test.rs +++ b/sgl-model-gateway/tests/routing/cache_aware_backward_compat_test.rs @@ -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> = 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> = 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> = 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> = 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"); }