[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:
Kan Wu
2026-09-20 16:57:07 -07:00
committed by GitHub
co-authored by Claude Fable 5.1
parent 4a9dc5c4af
commit aedda8377e
15 changed files with 2566 additions and 4 deletions
@@ -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)));
}
}
@@ -21,6 +21,8 @@ use std::sync::Arc;
use std::time::Duration;
use tower::ServiceExt;
mod reorg;
const TEST_TIMEOUT: Duration = Duration::from_secs(5);
fn config_for(_worker_url: &str) -> Config {
@@ -0,0 +1,523 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use super::*;
use crate::common::mock_worker::MockWorker;
use futures::future::BoxFuture;
use sgl_router::buckets_reorg::{Bucket, BucketGroups, BucketResolver, EngineGroup};
use sgl_router::policies::PolicyRegistry;
use sgl_router::policies_reorg::admission::{AllowAll, Decision, EngineAdmission};
use sgl_router::policies_reorg::{Pick, PickError, PickRequest, Policy, Rejection, Stage};
use sgl_router::server::app_context::ChatRouting;
use sgl_router::state::load_monitor::engine_reported_load::EngineReportedWorkerLoad;
use std::sync::Mutex;
type PickCall = (String, Stage, u64, Option<u64>);
#[derive(Debug)]
struct FirstPolicy {
calls: Mutex<Vec<PickCall>>,
admission: Arc<dyn EngineAdmission>,
miss: bool,
invalid: bool,
}
impl Default for FirstPolicy {
fn default() -> Self {
Self {
admission: Arc::new(AllowAll),
miss: false,
invalid: false,
calls: Mutex::new(Vec::new()),
}
}
}
impl Policy for FirstPolicy {
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(),
request.stage,
request.input_tokens,
request.expected_peak_tokens,
));
if self.invalid {
return Err(PickError::InvalidSignal("invalid policy input".into()));
}
if self.miss {
return Err(PickError::NoCandidates);
}
if engines.is_empty() {
return Err(PickError::NoCandidates);
}
let engine = 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 RejectAll;
impl EngineAdmission for RejectAll {
fn check(
&self,
_: &Worker,
_: &PickRequest<'_>,
_: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
Ok(Decision::Reject("full".into()))
}
}
fn rejecting_policy() -> Arc<FirstPolicy> {
Arc::new(FirstPolicy {
admission: Arc::new(RejectAll),
..Default::default()
})
}
fn group(id: &str, policy: Arc<FirstPolicy>) -> EngineGroup {
EngineGroup {
worker_ids: Some([WorkerId(id.into())].into_iter().collect()),
policy,
}
}
fn context(workers: &[(&str, Stage, &MockWorker)], buckets: Vec<Bucket>) -> Arc<AppContext> {
let config = config_for("");
let registry = Arc::new(WorkerRegistry::default());
for &(id, mode, worker) in workers {
registry
.add(WorkerSpec {
id: WorkerId(id.into()),
url: worker.url.clone(),
mode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: Some(8998),
})
.unwrap();
}
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&config).unwrap());
let mut ctx = AppContext::new(
config,
tokenizers,
Arc::new(Proxy::new(TEST_TIMEOUT).unwrap()),
registry,
// The reorg route must not require the legacy policy registry.
Arc::new(PolicyRegistry::default()),
);
ctx.chat_routing = ChatRouting::Reorg(
[(ModelId("tiny".into()), BucketResolver::new(buckets))]
.into_iter()
.collect(),
);
Arc::new(ctx)
}
fn request(value: serde_json::Value) -> Request<Body> {
Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&value).unwrap()))
.unwrap()
}
fn body(content: &str) -> serde_json::Value {
serde_json::json!({"model": "tiny", "messages": [{"role": "user", "content": content}]})
}
#[tokio::test]
async fn length_selects_plain_bucket_before_engine_selection() {
let short_worker = MockWorker::start(vec![]).await;
let long_worker = MockWorker::start(vec![]).await;
let policy = Arc::new(FirstPolicy::default());
let mut short = Bucket::new("short", BucketGroups::Plain(group("short", policy.clone())));
short.limits.max = Some(4);
let long = Bucket::new("long", BucketGroups::Plain(group("long", policy.clone())));
let ctx = context(
&[
("short", Stage::Plain, &short_worker),
("long", Stage::Plain, &long_worker),
],
vec![long, short],
);
let app = build_router(ctx);
let response = app.clone().oneshot(request(body("hi"))).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let _ = response.into_body().collect().await.unwrap();
assert!(short_worker.captured.lock().unwrap().last_body.is_some());
assert!(long_worker.captured.lock().unwrap().last_body.is_none());
let response = app
.oneshot(request(body(&"hello ".repeat(30))))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let _ = response.into_body().collect().await.unwrap();
assert!(long_worker.captured.lock().unwrap().last_body.is_some());
let calls = policy.calls.lock().unwrap();
assert_eq!(calls.len(), 2);
assert_eq!((&*calls[0].0, calls[0].1), ("short", Stage::Plain));
assert_eq!((&*calls[1].0, calls[1].1), ("long", Stage::Plain));
}
#[tokio::test]
async fn pd_picks_both_groups_from_selected_bucket_and_shares_bootstrap() {
let prefill = MockWorker::start(vec![]).await;
let decode = MockWorker::start(vec![]).await;
let policy = Arc::new(FirstPolicy::default());
let mut selected = Bucket::new(
"selected",
BucketGroups::Pd {
prefill: group("p", policy.clone()),
decode: group("d", policy.clone()),
},
);
selected.limits.max = Some(100);
let other_policy = Arc::new(FirstPolicy::default());
let other = Bucket::new(
"other",
BucketGroups::Pd {
prefill: group("p", other_policy.clone()),
decode: group("d", other_policy.clone()),
},
);
let ctx = context(
&[
("p", Stage::Prefill, &prefill),
("d", Stage::Decode, &decode),
],
vec![other, selected],
);
let mut body = body("hello");
body["max_completion_tokens"] = 10.into();
let response = build_router(ctx).oneshot(request(body)).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()["x-sgl-decode-url"], decode.url);
let _ = response.into_body().collect().await.unwrap();
tokio::time::timeout(TEST_TIMEOUT, async {
while prefill.captured.lock().unwrap().last_body.is_none() {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
let p: serde_json::Value =
serde_json::from_slice(prefill.captured.lock().unwrap().last_body.as_ref().unwrap())
.unwrap();
let d: serde_json::Value =
serde_json::from_slice(decode.captured.lock().unwrap().last_body.as_ref().unwrap())
.unwrap();
assert!(p["bootstrap_room"].is_number());
assert_eq!(p["bootstrap_room"], d["bootstrap_room"]);
let calls = policy.calls.lock().unwrap();
assert_eq!(calls.len(), 2);
assert_eq!((&*calls[0].0, calls[0].1), ("selected", Stage::Prefill));
assert_eq!((&*calls[1].0, calls[1].1), ("selected", Stage::Decode));
assert_eq!(calls[0].3, Some(calls[0].2 + 10));
assert_eq!(calls[1].3, calls[0].3);
assert!(other_policy.calls.lock().unwrap().is_empty());
}
#[tokio::test]
async fn missing_decode_in_all_buckets_does_not_dispatch_prefill() {
let prefill = MockWorker::start(vec![]).await;
let decode = MockWorker::start(vec![]).await;
let policy = Arc::new(FirstPolicy::default());
let mut selected = Bucket::new(
"selected",
BucketGroups::Pd {
prefill: group("p", policy.clone()),
decode: group("missing", policy.clone()),
},
);
selected.limits.max = Some(100);
let other = Bucket::new(
"other",
BucketGroups::Pd {
prefill: group("p", policy.clone()),
decode: group("also-missing", policy.clone()),
},
);
let ctx = context(
&[
("p", Stage::Prefill, &prefill),
("d", Stage::Decode, &decode),
],
vec![other, selected],
);
let response = build_router(ctx.clone())
.oneshot(request(body("hi")))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response.headers()["x-router-error-code"],
"no_decode_workers_available"
);
assert!(prefill.captured.lock().unwrap().last_body.is_none());
assert!(decode.captured.lock().unwrap().last_body.is_none());
assert_eq!(policy.calls.lock().unwrap().len(), 2);
assert_eq!(ctx.router_inflight_load.inflight_count(), 0);
assert_eq!(
ctx.registry
.get(&WorkerId("p".into()))
.unwrap()
.router_inflight_load(),
0
);
}
#[tokio::test]
async fn rejects_unsupported_length_unknown_model_and_overflow_before_policy() {
let worker = MockWorker::start(vec![]).await;
let policy = Arc::new(FirstPolicy::default());
let mut bucket = Bucket::new("short", BucketGroups::Plain(group("w", policy.clone())));
bucket.max_context_tokens = Some(4);
let app = build_router(context(&[("w", Stage::Plain, &worker)], vec![bucket]));
let mut long = body("hi");
long["max_tokens"] = 100.into();
let mut overflow = body("hi");
overflow["max_tokens"] = u64::MAX.into();
let mut unknown = body("hi");
unknown["model"] = "unknown".into();
for (body, status) in [
(long, StatusCode::BAD_REQUEST),
(overflow, StatusCode::BAD_REQUEST),
(unknown, StatusCode::NOT_FOUND),
(serde_json::json!({}), StatusCode::BAD_REQUEST),
] {
let response = app.clone().oneshot(request(body)).await.unwrap();
assert_eq!(response.status(), status);
}
assert!(policy.calls.lock().unwrap().is_empty());
assert!(worker.captured.lock().unwrap().last_body.is_none());
}
#[tokio::test]
async fn streaming_uses_existing_forwarder() {
let worker = MockWorker::start(vec!["data: {\"choices\":[]}\n\n", "data: [DONE]\n\n"]).await;
let policy = Arc::new(FirstPolicy::default());
let bucket = Bucket::new("plain", BucketGroups::Plain(group("w", policy)));
let app = build_router(context(&[("w", Stage::Plain, &worker)], vec![bucket]));
let mut body = body("hi");
body["stream"] = true.into();
let response = app.oneshot(request(body)).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers()["content-type"]
.to_str()
.unwrap()
.starts_with("text/event-stream"));
let bytes = response.into_body().collect().await.unwrap().to_bytes();
assert!(String::from_utf8_lossy(&bytes).contains("data: [DONE]"));
}
#[tokio::test]
async fn reorg_route_keeps_chat_body_limit() {
let app = build_router(context(&[], vec![]));
let request = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(vec![b' '; MAX_CHAT_BODY_BYTES + 1]))
.unwrap();
assert_eq!(
app.oneshot(request).await.unwrap().status(),
StatusCode::PAYLOAD_TOO_LARGE
);
}
#[tokio::test]
async fn plain_fallback_skips_empty_missed_and_rejected_buckets_then_stops_on_success() {
let rejected_worker = MockWorker::start(vec![]).await;
let winner = MockWorker::start(vec![]).await;
let skipped = Arc::new(FirstPolicy::default());
let rejected = rejecting_policy();
let missed = Arc::new(FirstPolicy {
miss: true,
..Default::default()
});
let accepted = Arc::new(FirstPolicy::default());
let buckets = vec![
Bucket::new(
"a-empty",
BucketGroups::Plain(group("missing", skipped.clone())),
),
Bucket::new(
"b-miss",
BucketGroups::Plain(group("rejected", missed.clone())),
),
Bucket::new(
"c-rejected",
BucketGroups::Plain(group("rejected", rejected.clone())),
),
Bucket::new(
"d-winner",
BucketGroups::Plain(group("winner", accepted.clone())),
),
Bucket::new(
"e-unused",
BucketGroups::Plain(group("winner", skipped.clone())),
),
];
let ctx = context(
&[
("rejected", Stage::Plain, &rejected_worker),
("winner", Stage::Plain, &winner),
],
buckets,
);
let response = build_router(ctx)
.oneshot(request(body("hi")))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let _ = response.into_body().collect().await.unwrap();
assert!(rejected_worker.captured.lock().unwrap().last_body.is_none());
assert!(winner.captured.lock().unwrap().last_body.is_some());
assert!(skipped.calls.lock().unwrap().is_empty());
assert_eq!(missed.calls.lock().unwrap().len(), 1);
assert_eq!(rejected.calls.lock().unwrap().len(), 1);
assert_eq!(accepted.calls.lock().unwrap()[0].0, "d-winner");
}
#[tokio::test]
async fn decode_failure_retries_both_groups_in_next_bucket_without_dispatching_first_prefill() {
// Cover an empty decode group and rejection of a selected decode engine.
for reject_decode in [false, true] {
let first_prefill = MockWorker::start(vec![]).await;
let second_prefill = MockWorker::start(vec![]).await;
let decode = MockWorker::start(vec![]).await;
let first = Arc::new(FirstPolicy::default());
let rejected = rejecting_policy();
let accepted = Arc::new(FirstPolicy::default());
let buckets = vec![
Bucket::new(
"a-first",
BucketGroups::Pd {
prefill: group("p1", first.clone()),
decode: group(if reject_decode { "d" } else { "missing" }, rejected),
},
),
Bucket::new(
"b-second",
BucketGroups::Pd {
prefill: group("p2", accepted.clone()),
decode: group("d", accepted.clone()),
},
),
];
let ctx = context(
&[
("p1", Stage::Prefill, &first_prefill),
("p2", Stage::Prefill, &second_prefill),
("d", Stage::Decode, &decode),
],
buckets,
);
let response = build_router(ctx.clone())
.oneshot(request(body("hi")))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let _ = response.into_body().collect().await.unwrap();
tokio::time::timeout(TEST_TIMEOUT, async {
while second_prefill.captured.lock().unwrap().last_body.is_none() {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
assert!(first_prefill.captured.lock().unwrap().last_body.is_none());
assert!(decode.captured.lock().unwrap().last_body.is_some());
assert_eq!(
ctx.registry
.get(&WorkerId("p1".into()))
.unwrap()
.router_inflight_load(),
0
);
assert_eq!(first.calls.lock().unwrap().len(), 1);
let calls = accepted.calls.lock().unwrap();
assert_eq!(calls.len(), 2);
assert_eq!((&*calls[0].0, calls[0].1), ("b-second", Stage::Prefill));
assert_eq!((&*calls[1].0, calls[1].1), ("b-second", Stage::Decode));
}
}
#[tokio::test]
async fn admission_exhaustion_is_preserved_when_later_buckets_are_empty() {
let worker = MockWorker::start(vec![]).await;
let first = rejecting_policy();
let second = rejecting_policy();
let empty = Arc::new(FirstPolicy::default());
let ctx = context(
&[("w", Stage::Plain, &worker)],
vec![
Bucket::new("a-first", BucketGroups::Plain(group("w", first.clone()))),
Bucket::new("b-second", BucketGroups::Plain(group("w", second.clone()))),
Bucket::new(
"c-empty",
BucketGroups::Plain(group("missing", empty.clone())),
),
],
);
let response = build_router(ctx.clone())
.oneshot(request(body("hi")))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response.headers()["x-router-error-code"],
"policy_selection_failed"
);
assert_eq!(first.calls.lock().unwrap().len(), 1);
assert_eq!(second.calls.lock().unwrap().len(), 1);
assert!(empty.calls.lock().unwrap().is_empty());
assert!(worker.captured.lock().unwrap().last_body.is_none());
assert_eq!(ctx.router_inflight_load.inflight_count(), 0);
}
#[tokio::test]
async fn invalid_policy_signal_stops_bucket_iteration() {
let worker = MockWorker::start(vec![]).await;
let invalid = Arc::new(FirstPolicy {
invalid: true,
..Default::default()
});
let later = Arc::new(FirstPolicy::default());
let ctx = context(
&[("w", Stage::Plain, &worker)],
vec![
Bucket::new(
"a-invalid",
BucketGroups::Plain(group("w", invalid.clone())),
),
Bucket::new("b-later", BucketGroups::Plain(group("w", later.clone()))),
],
);
let response = build_router(ctx)
.oneshot(request(body("hi")))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(invalid.calls.lock().unwrap().len(), 1);
assert!(later.calls.lock().unwrap().is_empty());
assert!(worker.captured.lock().unwrap().last_body.is_none());
}