[sgl-router] refactor - layout BucketResolver, Bucket, EngineGroup and implement PowerOfTwo (#40241)
Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5.1
parent
4a9dc5c4af
commit
aedda8377e
@@ -10,5 +10,8 @@
|
||||
mod discovery;
|
||||
mod health;
|
||||
mod policies;
|
||||
mod policies_reorg;
|
||||
mod policies_reorg_load;
|
||||
mod policies_reorg_power_of_two;
|
||||
mod tokenizer;
|
||||
mod workers;
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures::future::BoxFuture;
|
||||
use sgl_router::buckets_reorg::{
|
||||
Bucket, BucketGroups, BucketRequest, BucketResolver, EngineGroup, TokenLimits,
|
||||
};
|
||||
use sgl_router::discovery::{ModelId, WorkerId, WorkerSpec};
|
||||
use sgl_router::policies_reorg::admission::{AllowAll, Decision, EngineAdmission};
|
||||
use sgl_router::policies_reorg::{Pick, PickError, PickRequest, Policy, Rejection, Stage};
|
||||
use sgl_router::state::load_monitor::engine_reported_load::EngineReportedWorkerLoad;
|
||||
use sgl_router::workers::{Worker, WorkerRegistry};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct TestPolicy {
|
||||
admission: Arc<dyn EngineAdmission>,
|
||||
result: Option<Arc<Worker>>,
|
||||
miss: bool,
|
||||
invalid: bool,
|
||||
calls: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl Default for TestPolicy {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
admission: Arc::new(AllowAll),
|
||||
result: None,
|
||||
miss: false,
|
||||
invalid: false,
|
||||
calls: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Policy for TestPolicy {
|
||||
fn pick<'a>(
|
||||
&'a self,
|
||||
engines: &'a [Arc<Worker>],
|
||||
request: &'a PickRequest<'a>,
|
||||
) -> BoxFuture<'a, Result<Pick, PickError>> {
|
||||
Box::pin(async move {
|
||||
self.calls.lock().unwrap().push(request.bucket.to_owned());
|
||||
if self.invalid {
|
||||
return Err(PickError::InvalidSignal("test signal".into()));
|
||||
}
|
||||
if self.miss {
|
||||
return Err(PickError::NoCandidates);
|
||||
}
|
||||
if engines.is_empty() {
|
||||
return Err(PickError::NoCandidates);
|
||||
}
|
||||
let engine = self.result.clone().unwrap_or_else(|| engines[0].clone());
|
||||
if let Decision::Reject(reason) = self.admission.check(&engine, request, None)? {
|
||||
return Err(PickError::AdmissionRejected(Rejection {
|
||||
engine: engine.id.clone(),
|
||||
reason,
|
||||
}));
|
||||
}
|
||||
Ok(Pick {
|
||||
engine,
|
||||
reason: "test",
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct Reject(&'static str);
|
||||
|
||||
impl EngineAdmission for Reject {
|
||||
fn check(
|
||||
&self,
|
||||
engine: &Worker,
|
||||
_: &PickRequest<'_>,
|
||||
_: Option<&EngineReportedWorkerLoad>,
|
||||
) -> Result<Decision, PickError> {
|
||||
Ok(if engine.id.0 == self.0 {
|
||||
Decision::Reject("full".into())
|
||||
} else {
|
||||
Decision::Allow
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn spec(id: &str, mode: Stage, model: &str) -> WorkerSpec {
|
||||
WorkerSpec {
|
||||
id: WorkerId(id.into()),
|
||||
url: format!("http://{id}"),
|
||||
mode,
|
||||
model_ids: vec![ModelId(model.into())],
|
||||
bootstrap_port: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn registry() -> Arc<WorkerRegistry> {
|
||||
let workers = Arc::new(WorkerRegistry::default());
|
||||
for (id, mode, model) in [
|
||||
("a", Stage::Plain, "m"),
|
||||
("b", Stage::Plain, "m"),
|
||||
("unhealthy", Stage::Plain, "m"),
|
||||
("other", Stage::Plain, "other"),
|
||||
("p", Stage::Prefill, "pd"),
|
||||
("d", Stage::Decode, "pd"),
|
||||
] {
|
||||
workers.add(spec(id, mode, model)).unwrap();
|
||||
}
|
||||
let unhealthy = workers.get(&WorkerId("unhealthy".into())).unwrap();
|
||||
for _ in 0..3 {
|
||||
unhealthy.breaker.record_failure();
|
||||
}
|
||||
workers
|
||||
}
|
||||
|
||||
fn group(members: &[&str], policy: Arc<dyn Policy>) -> EngineGroup {
|
||||
EngineGroup {
|
||||
worker_ids: Some(members.iter().map(|id| WorkerId((*id).into())).collect()),
|
||||
policy,
|
||||
}
|
||||
}
|
||||
|
||||
fn bucket(id: &str, max: Option<u64>, policy: Arc<dyn Policy>) -> Bucket {
|
||||
let mut bucket = Bucket::new(id, BucketGroups::Plain(EngineGroup::new(policy)));
|
||||
bucket.limits.max = max;
|
||||
bucket
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn groups_isolate_model_health_stage_and_membership() {
|
||||
let workers = registry();
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let group = EngineGroup::new(Arc::new(TestPolicy::default()));
|
||||
assert_eq!(
|
||||
group.pick(&workers, &request).await.unwrap().engine.id.0,
|
||||
"a"
|
||||
);
|
||||
workers.remove(&WorkerId("a".into()));
|
||||
assert_eq!(
|
||||
group.pick(&workers, &request).await.unwrap().engine.id.0,
|
||||
"b"
|
||||
);
|
||||
workers.remove(&WorkerId("b".into()));
|
||||
assert!(matches!(
|
||||
group.pick(&workers, &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
|
||||
let pd = ModelId("pd".into());
|
||||
for (stage, expected) in [(Stage::Prefill, "p"), (Stage::Decode, "d")] {
|
||||
let request = PickRequest::new(&pd, stage, 10);
|
||||
let group = self::group(
|
||||
&["p", "d", "other", "unhealthy"],
|
||||
Arc::new(TestPolicy::default()),
|
||||
);
|
||||
assert_eq!(
|
||||
group.pick(&workers, &request).await.unwrap().engine.id.0,
|
||||
expected
|
||||
);
|
||||
}
|
||||
let empty = self::group(&[], Arc::new(TestPolicy::default()));
|
||||
assert!(matches!(
|
||||
empty.pick(&workers, &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_orders_all_length_fits_by_capacity_rank_and_id() {
|
||||
let policy = Arc::new(TestPolicy::default());
|
||||
let mut z = bucket("z", Some(20), policy.clone());
|
||||
z.rank = 1;
|
||||
let mut a = bucket("a", Some(20), policy.clone());
|
||||
a.rank = 1;
|
||||
let mut later = bucket("later", Some(20), policy.clone());
|
||||
later.rank = 2;
|
||||
let mut min = bucket("min", Some(15), policy.clone());
|
||||
min.limits.min = Some(11);
|
||||
let resolver = BucketResolver::new(vec![
|
||||
bucket("catch-all", None, policy.clone()),
|
||||
bucket("too-small", Some(9), policy),
|
||||
z,
|
||||
later,
|
||||
a,
|
||||
min,
|
||||
]);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.resolve(10, None)
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|bucket| bucket.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["a", "z", "later", "catch-all"]
|
||||
);
|
||||
assert_eq!(resolver.resolve(11, None).unwrap()[0].id, "min");
|
||||
assert_eq!(resolver.resolve(15, None).unwrap()[0].id, "min");
|
||||
assert_eq!(resolver.resolve(20, None).unwrap()[0].id, "a");
|
||||
assert_eq!(resolver.resolve(21, None).unwrap()[0].id, "catch-all");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_capacity_checks_peak_when_known_and_input_otherwise() {
|
||||
let policy = Arc::new(TestPolicy::default());
|
||||
let mut short = bucket("short", None, policy.clone());
|
||||
short.max_context_tokens = Some(20);
|
||||
let mut long = bucket("long", None, policy);
|
||||
long.max_context_tokens = Some(30);
|
||||
let resolver = BucketResolver::new(vec![long, short]);
|
||||
assert_eq!(resolver.resolve(10, None).unwrap()[0].id, "short");
|
||||
assert_eq!(resolver.resolve(10, Some(20)).unwrap()[0].id, "short");
|
||||
assert_eq!(resolver.resolve(10, Some(21)).unwrap()[0].id, "long");
|
||||
assert!(resolver.resolve(10, Some(31)).unwrap().is_empty());
|
||||
assert!(resolver.resolve(31, None).unwrap().is_empty());
|
||||
assert!(matches!(
|
||||
resolver.resolve(10, Some(9)),
|
||||
Err(PickError::InvalidSignal(_))
|
||||
));
|
||||
assert!(BucketResolver::default()
|
||||
.resolve(1, None)
|
||||
.unwrap()
|
||||
.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selected_pd_bucket_owns_both_memberships_and_policies() {
|
||||
let workers = registry();
|
||||
workers.add(spec("p2", Stage::Prefill, "pd")).unwrap();
|
||||
workers.add(spec("d2", Stage::Decode, "pd")).unwrap();
|
||||
let model = ModelId("pd".into());
|
||||
let prefill_policy = Arc::new(TestPolicy::default());
|
||||
let decode_policy = Arc::new(TestPolicy::default());
|
||||
let resolver = BucketResolver::new(vec![Bucket::new(
|
||||
"shared",
|
||||
BucketGroups::Pd {
|
||||
prefill: group(&["p2", "d", "a"], prefill_policy.clone()),
|
||||
decode: group(&["d2", "p", "other"], decode_policy.clone()),
|
||||
},
|
||||
)]);
|
||||
let bucket = resolver.resolve(10, Some(20)).unwrap()[0];
|
||||
let request = BucketRequest {
|
||||
model: &model,
|
||||
input_tokens: 10,
|
||||
expected_peak_tokens: Some(20),
|
||||
token_ids: None,
|
||||
session_key: None,
|
||||
routing_key: None,
|
||||
};
|
||||
let picks = bucket.pick_engines(&workers, &request).await.unwrap();
|
||||
assert_eq!(picks.prefill.engine.id.0, "p2");
|
||||
assert_eq!(picks.decode.unwrap().engine.id.0, "d2");
|
||||
assert_eq!(*prefill_policy.calls.lock().unwrap(), ["shared"]);
|
||||
assert_eq!(*decode_policy.calls.lock().unwrap(), ["shared"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolver_includes_empty_groups_without_invoking_policies() {
|
||||
let workers = registry();
|
||||
let model = ModelId("m".into());
|
||||
let policy = Arc::new(TestPolicy::default());
|
||||
let mut empty = Bucket::new(
|
||||
"empty",
|
||||
BucketGroups::Plain(group(&["missing"], policy.clone())),
|
||||
);
|
||||
empty.limits = TokenLimits {
|
||||
min: None,
|
||||
max: Some(10),
|
||||
};
|
||||
let resolver = BucketResolver::new(vec![empty, bucket("available", Some(20), policy.clone())]);
|
||||
let buckets = resolver.resolve(10, None).unwrap();
|
||||
assert_eq!(
|
||||
buckets
|
||||
.iter()
|
||||
.map(|bucket| bucket.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["empty", "available"]
|
||||
);
|
||||
let bucket = buckets[0];
|
||||
let BucketGroups::Plain(group) = &bucket.groups else {
|
||||
panic!("expected plain")
|
||||
};
|
||||
let request = PickRequest {
|
||||
bucket: &bucket.id,
|
||||
..PickRequest::new(&model, Stage::Plain, 10)
|
||||
};
|
||||
assert!(matches!(
|
||||
group.pick(&workers, &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
assert!(policy.calls.lock().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn group_propagates_rejections_misses_and_invalid_signals() {
|
||||
let workers = registry();
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let rejected = group(
|
||||
&["a"],
|
||||
Arc::new(TestPolicy {
|
||||
admission: Arc::new(Reject("a")),
|
||||
..Default::default()
|
||||
}),
|
||||
);
|
||||
assert!(matches!(
|
||||
rejected.pick(&workers, &request).await,
|
||||
Err(PickError::AdmissionRejected(_))
|
||||
));
|
||||
let miss = group(
|
||||
&["a"],
|
||||
Arc::new(TestPolicy {
|
||||
miss: true,
|
||||
..Default::default()
|
||||
}),
|
||||
);
|
||||
assert!(matches!(
|
||||
miss.pick(&workers, &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
let invalid = group(
|
||||
&["a"],
|
||||
Arc::new(TestPolicy {
|
||||
invalid: true,
|
||||
..Default::default()
|
||||
}),
|
||||
);
|
||||
assert!(matches!(
|
||||
invalid.pick(&workers, &request).await,
|
||||
Err(PickError::InvalidSignal(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn group_rejects_foreign_pick_even_with_same_worker_id() {
|
||||
let workers = registry();
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let group = group(
|
||||
&["a"],
|
||||
Arc::new(TestPolicy {
|
||||
result: Some(Arc::new(Worker::new(spec("a", Stage::Plain, "m")))),
|
||||
..Default::default()
|
||||
}),
|
||||
);
|
||||
assert!(matches!(
|
||||
group.pick(&workers, &request).await,
|
||||
Err(PickError::OutsideCandidates(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selected_engine_rejection_does_not_try_an_alternative() {
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let engines = [
|
||||
Arc::new(Worker::new(spec("a", Stage::Plain, "m"))),
|
||||
Arc::new(Worker::new(spec("b", Stage::Plain, "m"))),
|
||||
];
|
||||
let policy = TestPolicy {
|
||||
admission: Arc::new(Reject("a")),
|
||||
..Default::default()
|
||||
};
|
||||
assert!(matches!(
|
||||
policy.pick(&engines, &request).await,
|
||||
Err(PickError::AdmissionRejected(_))
|
||||
));
|
||||
assert!(matches!(
|
||||
policy.pick(&[], &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bucket_scopes_plain_pick_and_preserves_request_facts() {
|
||||
#[derive(Debug)]
|
||||
struct InspectRequest;
|
||||
|
||||
impl Policy for InspectRequest {
|
||||
fn pick<'a>(
|
||||
&'a self,
|
||||
engines: &'a [Arc<Worker>],
|
||||
request: &'a PickRequest<'a>,
|
||||
) -> BoxFuture<'a, Result<Pick, PickError>> {
|
||||
Box::pin(async move {
|
||||
assert_eq!(request.model.0, "m");
|
||||
assert_eq!(request.bucket, "plain-bucket");
|
||||
assert_eq!(request.stage, Stage::Plain);
|
||||
assert_eq!(request.input_tokens, 2);
|
||||
assert_eq!(request.expected_peak_tokens, Some(12));
|
||||
assert_eq!(request.token_ids, Some([7, 9].as_slice()));
|
||||
assert_eq!(request.session_key, Some("session"));
|
||||
assert_eq!(request.routing_key, Some("routing"));
|
||||
Ok(Pick {
|
||||
engine: engines[0].clone(),
|
||||
reason: "inspected",
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
let workers = registry();
|
||||
let model = ModelId("m".into());
|
||||
let bucket = Bucket::new(
|
||||
"plain-bucket",
|
||||
BucketGroups::Plain(group(&["b"], Arc::new(InspectRequest))),
|
||||
);
|
||||
let request = BucketRequest {
|
||||
model: &model,
|
||||
input_tokens: 2,
|
||||
expected_peak_tokens: Some(12),
|
||||
token_ids: Some(&[7, 9]),
|
||||
session_key: Some("session"),
|
||||
routing_key: Some("routing"),
|
||||
};
|
||||
let picks = bucket.pick_engines(&workers, &request).await.unwrap();
|
||||
assert_eq!(picks.prefill.engine.id.0, "b");
|
||||
assert_eq!(picks.prefill.reason, "inspected");
|
||||
assert!(picks.decode.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn power_of_two_checks_selected_engine_and_propagates_rejection_without_fallback() {
|
||||
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
|
||||
use sgl_router::state::load_monitor::engine_reported_load::EngineReportedLoadTable;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct Check {
|
||||
calls: Mutex<Vec<WorkerId>>,
|
||||
reject: bool,
|
||||
invalid: bool,
|
||||
}
|
||||
|
||||
impl EngineAdmission for Check {
|
||||
fn check(
|
||||
&self,
|
||||
engine: &Worker,
|
||||
_: &PickRequest<'_>,
|
||||
_: Option<&EngineReportedWorkerLoad>,
|
||||
) -> Result<Decision, PickError> {
|
||||
self.calls.lock().unwrap().push(engine.id.clone());
|
||||
if self.invalid {
|
||||
Err(PickError::InvalidSignal("admission input".into()))
|
||||
} else if self.reject {
|
||||
Ok(Decision::Reject("full".into()))
|
||||
} else {
|
||||
Ok(Decision::Allow)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let engine = Arc::new(Worker::new(spec("a", Stage::Plain, "m")));
|
||||
let other = Arc::new(Worker::new(spec("b", Stage::Plain, "m")));
|
||||
let _busy = other.load_guard();
|
||||
for engines in [vec![engine.clone()], vec![other.clone(), engine.clone()]] {
|
||||
for (reject, invalid) in [(false, false), (true, false), (false, true)] {
|
||||
let check = Arc::new(Check {
|
||||
calls: Mutex::new(Vec::new()),
|
||||
reject,
|
||||
invalid,
|
||||
});
|
||||
let mut policy = PowerOfTwoPolicy::new(EngineReportedLoadTable::new());
|
||||
policy.admission = check.clone();
|
||||
assert!(matches!(
|
||||
policy.pick(&[], &request).await,
|
||||
Err(PickError::NoCandidates)
|
||||
));
|
||||
assert!(check.calls.lock().unwrap().is_empty());
|
||||
let result = policy.pick(&engines, &request).await;
|
||||
if invalid {
|
||||
assert!(matches!(result, Err(PickError::InvalidSignal(_))));
|
||||
} else if reject {
|
||||
assert!(matches!(result, Err(PickError::AdmissionRejected(reason))
|
||||
if reason.engine == engine.id && reason.reason == "full"));
|
||||
} else {
|
||||
assert!(Arc::ptr_eq(&result.unwrap().engine, &engine));
|
||||
}
|
||||
assert_eq!(
|
||||
check.calls.lock().unwrap().as_slice(),
|
||||
std::slice::from_ref(&engine.id)
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use sgl_router::discovery::{ModelId, WorkerId, WorkerSpec};
|
||||
use sgl_router::policies_reorg::admission::{Decision, EngineAdmission};
|
||||
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
|
||||
use sgl_router::policies_reorg::{PickError, PickRequest, Policy, Stage};
|
||||
use sgl_router::state::load_monitor::engine_reported_load::{
|
||||
EngineReportedLoadTable, EngineReportedWorkerLoad, LoadStat,
|
||||
};
|
||||
use sgl_router::workers::Worker;
|
||||
|
||||
const URL: &str = "http://engine";
|
||||
|
||||
fn report(table: &EngineReportedLoadTable, rank: u32, running: u64, at: Instant) {
|
||||
table.set(
|
||||
URL,
|
||||
rank,
|
||||
LoadStat {
|
||||
num_running_reqs: running,
|
||||
num_waiting_reqs: 2,
|
||||
num_tokens: 30,
|
||||
max_total_num_tokens: 100,
|
||||
native_cache: None,
|
||||
},
|
||||
at,
|
||||
);
|
||||
}
|
||||
|
||||
fn engine() -> Arc<Worker> {
|
||||
Arc::new(Worker::new(WorkerSpec {
|
||||
id: WorkerId("a".into()),
|
||||
url: URL.into(),
|
||||
mode: Stage::Plain,
|
||||
model_ids: vec![ModelId("m".into())],
|
||||
bootstrap_port: None,
|
||||
}))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ObserveAdmission {
|
||||
table: Arc<EngineReportedLoadTable>,
|
||||
observations: Mutex<Vec<Option<EngineReportedWorkerLoad>>>,
|
||||
}
|
||||
|
||||
impl EngineAdmission for ObserveAdmission {
|
||||
fn check(
|
||||
&self,
|
||||
engine: &Worker,
|
||||
_: &PickRequest<'_>,
|
||||
load: Option<&EngineReportedWorkerLoad>,
|
||||
) -> Result<Decision, PickError> {
|
||||
assert_eq!(engine.url, URL);
|
||||
// A new report arriving after selection must not change the observation
|
||||
// supplied to admission. The next pick should read the new report.
|
||||
report(&self.table, 0, 99, Instant::now());
|
||||
self.observations.lock().unwrap().push(load.cloned());
|
||||
Ok(Decision::Allow)
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selected_load_reaches_admission_and_next_pick_reads_fresh_state() {
|
||||
let table = EngineReportedLoadTable::new();
|
||||
let first_at = Instant::now();
|
||||
report(&table, 0, 1, first_at);
|
||||
report(&table, 1, 3, first_at);
|
||||
table.mark_expected_rank(URL, 0);
|
||||
table.mark_expected_rank(URL, 1);
|
||||
let admission = Arc::new(ObserveAdmission {
|
||||
table: table.clone(),
|
||||
observations: Mutex::default(),
|
||||
});
|
||||
let mut policy = PowerOfTwoPolicy::new(table);
|
||||
policy.admission = admission.clone();
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
let alternative = Arc::new(Worker::new(WorkerSpec {
|
||||
id: WorkerId("b".into()),
|
||||
url: "http://other".into(),
|
||||
mode: Stage::Plain,
|
||||
model_ids: vec![ModelId("m".into())],
|
||||
bootstrap_port: None,
|
||||
}));
|
||||
// These old-format reports lack native pressure metrics, so selection uses
|
||||
// local active counts for both candidates and chooses the second engine.
|
||||
let _busy = alternative.load_guard();
|
||||
let engines = [alternative, engine()];
|
||||
for _ in 0..2 {
|
||||
let pick = policy.pick(&engines, &request).await.unwrap();
|
||||
assert!(Arc::ptr_eq(&pick.engine, &engines[1]));
|
||||
}
|
||||
let observations = admission.observations.lock().unwrap();
|
||||
assert_eq!(observations.len(), 2);
|
||||
assert_eq!(
|
||||
observations[0],
|
||||
Some(EngineReportedWorkerLoad {
|
||||
num_running_reqs: 4,
|
||||
num_waiting_reqs: 4,
|
||||
num_tokens: 60,
|
||||
max_total_num_tokens: 200,
|
||||
captured_at: first_at,
|
||||
})
|
||||
);
|
||||
assert_eq!(observations[1].as_ref().unwrap().num_running_reqs, 102);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn missing_stale_and_incomplete_reports_reach_admission_as_unknown() {
|
||||
for case in ["missing", "stale", "incomplete"] {
|
||||
let table = EngineReportedLoadTable::new();
|
||||
match case {
|
||||
"missing" => {}
|
||||
"stale" => report(&table, 0, 1, Instant::now() - Duration::from_secs(3600)),
|
||||
"incomplete" => {
|
||||
report(&table, 0, 1, Instant::now());
|
||||
table.mark_expected_rank(URL, 0);
|
||||
table.mark_expected_rank(URL, 1);
|
||||
}
|
||||
_ => unreachable!(),
|
||||
}
|
||||
let admission = Arc::new(ObserveAdmission {
|
||||
table: table.clone(),
|
||||
observations: Mutex::default(),
|
||||
});
|
||||
let mut policy = PowerOfTwoPolicy::new(table);
|
||||
policy.admission = admission.clone();
|
||||
let model = ModelId("m".into());
|
||||
policy
|
||||
.pick(&[engine()], &PickRequest::new(&model, Stage::Plain, 10))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(*admission.observations.lock().unwrap(), [None], "{case}");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use sgl_router::discovery::{ModelId, WorkerId, WorkerSpec};
|
||||
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
|
||||
use sgl_router::policies_reorg::{PickRequest, Policy, Stage};
|
||||
use sgl_router::state::load_monitor::engine_reported_load::{
|
||||
EngineReportedLoadTable, LoadStat, NativeCacheRankLoad,
|
||||
};
|
||||
use sgl_router::workers::Worker;
|
||||
|
||||
fn engine(id: &str, stage: Stage, active: usize) -> Arc<Worker> {
|
||||
let worker = Arc::new(Worker::new(WorkerSpec {
|
||||
id: WorkerId(id.into()),
|
||||
url: format!("http://{id}"),
|
||||
mode: stage,
|
||||
model_ids: vec![ModelId("m".into())],
|
||||
bootstrap_port: None,
|
||||
}));
|
||||
worker.active_requests.store(active, Ordering::Relaxed);
|
||||
worker
|
||||
}
|
||||
|
||||
fn load(running: u64, waiting: u64, tokens: u64, capacity: u64, pending: u64) -> LoadStat {
|
||||
LoadStat {
|
||||
num_running_reqs: running,
|
||||
num_waiting_reqs: waiting,
|
||||
num_tokens: tokens,
|
||||
max_total_num_tokens: capacity,
|
||||
native_cache: Some(NativeCacheRankLoad {
|
||||
num_waiting_uncached_tokens: pending,
|
||||
num_total_tokens: tokens,
|
||||
max_running_requests: 100,
|
||||
total_prefill_uncached_tokens: 0,
|
||||
total_prefill_busy_us: 0,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
async fn assert_winner(
|
||||
policy: &PowerOfTwoPolicy,
|
||||
engines: &[Arc<Worker>],
|
||||
stage: Stage,
|
||||
expected: &Arc<Worker>,
|
||||
) {
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, stage, 10);
|
||||
// With exactly two candidates the winner is independent of sample order.
|
||||
for _ in 0..16 {
|
||||
let pick = policy.pick(engines, &request).await.unwrap();
|
||||
assert!(Arc::ptr_eq(&pick.engine, expected), "stage: {stage:?}");
|
||||
assert_eq!(pick.reason, "power_of_two");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn plain_and_prefill_use_pending_work_while_decode_uses_request_pressure() {
|
||||
for stage in [Stage::Plain, Stage::Prefill, Stage::Decode] {
|
||||
let engines = [engine("a", stage, 100), engine("b", stage, 0)];
|
||||
let table = EngineReportedLoadTable::new();
|
||||
table.set(&engines[0].url, 0, load(5, 8, 10, 100, 1), Instant::now());
|
||||
table.set(&engines[1].url, 0, load(1, 1, 10, 100, 100), Instant::now());
|
||||
let expected = if stage == Stage::Decode { 1 } else { 0 };
|
||||
assert_winner(
|
||||
&PowerOfTwoPolicy::new(table),
|
||||
&engines,
|
||||
stage,
|
||||
&engines[expected],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn prefill_uses_estimated_queue_time_only_when_both_engines_have_rates() {
|
||||
for both_have_rates in [true, false] {
|
||||
let engines = [
|
||||
engine("a", Stage::Prefill, 0),
|
||||
engine("b", Stage::Prefill, 0),
|
||||
];
|
||||
let table = EngineReportedLoadTable::new();
|
||||
for (i, worker) in engines.iter().enumerate() {
|
||||
let mut report = load(1, 1, 10, 100, if i == 0 { 10 } else { 20 });
|
||||
if i == 0 || both_have_rates {
|
||||
table.set(&worker.url, 0, report.clone(), Instant::now());
|
||||
}
|
||||
let native = report.native_cache.as_mut().unwrap();
|
||||
native.total_prefill_uncached_tokens = if i == 0 { 100 } else { 1000 };
|
||||
native.total_prefill_busy_us = 1_000_000;
|
||||
table.set(&worker.url, 0, report, Instant::now());
|
||||
}
|
||||
// B has more queued tokens, but its higher throughput gives a shorter
|
||||
// estimated queue. Without B's rate, compare queued tokens for both.
|
||||
let expected = usize::from(both_have_rates);
|
||||
assert_winner(
|
||||
&PowerOfTwoPolicy::new(table),
|
||||
&engines,
|
||||
Stage::Prefill,
|
||||
&engines[expected],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn decode_orders_by_waiting_running_kv_fraction_then_tokens() {
|
||||
let cases = [
|
||||
(load(50, 1, 90, 100, 0), load(1, 2, 1, 100, 0)),
|
||||
(load(1, 1, 90, 100, 0), load(2, 1, 1, 100, 0)),
|
||||
(load(1, 1, 100, 1000, 0), load(1, 1, 20, 100, 0)),
|
||||
(load(1, 1, 10, 100, 0), load(1, 1, 100, 1000, 0)),
|
||||
];
|
||||
for (left, right) in cases {
|
||||
let engines = [
|
||||
engine("a", Stage::Decode, 100),
|
||||
engine("b", Stage::Decode, 0),
|
||||
];
|
||||
let table = EngineReportedLoadTable::new();
|
||||
table.set(&engines[0].url, 0, left, Instant::now());
|
||||
table.set(&engines[1].url, 0, right, Instant::now());
|
||||
assert_winner(
|
||||
&PowerOfTwoPolicy::new(table),
|
||||
&engines,
|
||||
Stage::Decode,
|
||||
&engines[0],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unusable_telemetry_falls_back_to_local_load_for_both_candidates() {
|
||||
for stage in [Stage::Plain, Stage::Prefill, Stage::Decode] {
|
||||
for case in [
|
||||
"missing",
|
||||
"stale",
|
||||
"incomplete",
|
||||
"old_publisher",
|
||||
"unknown_capacity",
|
||||
] {
|
||||
let engines = [engine("a", stage, 1), engine("b", stage, 5)];
|
||||
let table = EngineReportedLoadTable::new();
|
||||
// A's high reported pressure must not be compared to B's local
|
||||
// count or to a fabricated zero for its unavailable telemetry.
|
||||
table.set(
|
||||
&engines[0].url,
|
||||
0,
|
||||
load(90, 90, 90, 100, 900),
|
||||
Instant::now(),
|
||||
);
|
||||
let mut right = load(0, 0, 0, 100, 0);
|
||||
let mut at = Instant::now();
|
||||
match case {
|
||||
"missing" => {}
|
||||
"stale" => at -= Duration::from_secs(3600),
|
||||
"incomplete" => {
|
||||
table.mark_expected_rank(&engines[1].url, 0);
|
||||
table.mark_expected_rank(&engines[1].url, 1);
|
||||
}
|
||||
"old_publisher" => right.native_cache = None,
|
||||
"unknown_capacity" => right.max_total_num_tokens = 0,
|
||||
_ => unreachable!(),
|
||||
}
|
||||
if case != "missing" {
|
||||
table.set(&engines[1].url, 0, right, at);
|
||||
}
|
||||
assert_winner(&PowerOfTwoPolicy::new(table), &engines, stage, &engines[0]).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn equal_reported_pressure_uses_local_active_load_as_tiebreaker() {
|
||||
for stage in [Stage::Plain, Stage::Prefill, Stage::Decode] {
|
||||
let engines = [engine("a", stage, 5), engine("b", stage, 1)];
|
||||
let table = EngineReportedLoadTable::new();
|
||||
for worker in &engines {
|
||||
table.set(&worker.url, 0, load(1, 1, 10, 100, 10), Instant::now());
|
||||
}
|
||||
assert_winner(&PowerOfTwoPolicy::new(table), &engines, stage, &engines[1]).await;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn multiple_candidates_never_select_the_unique_busiest_engine() {
|
||||
let engines: Vec<_> = (0..8)
|
||||
.map(|i| engine(&i.to_string(), Stage::Plain, i))
|
||||
.collect();
|
||||
let policy = PowerOfTwoPolicy::new(EngineReportedLoadTable::new());
|
||||
let model = ModelId("m".into());
|
||||
let request = PickRequest::new(&model, Stage::Plain, 10);
|
||||
for _ in 0..64 {
|
||||
let pick = policy.pick(&engines, &request).await.unwrap();
|
||||
// Every distinct pair has an engine less busy than the last candidate.
|
||||
assert!(engines[..7]
|
||||
.iter()
|
||||
.any(|engine| Arc::ptr_eq(engine, &pick.engine)));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user