[sgl-router] refactor - SLO ordering for bucket selection (#40292)

This commit is contained in:
Kan Wu
2026-09-22 11:16:10 +08:00
committed by GitHub
parent 27f796ca6c
commit 59a723ef1e
8 changed files with 422 additions and 37 deletions
@@ -16,5 +16,6 @@ mod policies_reorg_cache_aware;
mod policies_reorg_load;
mod policies_reorg_power_of_two;
mod policies_reorg_session_aware;
mod policies_reorg_slo;
mod tokenizer;
mod workers;
@@ -186,17 +186,20 @@ fn resolve_orders_all_length_fits_by_capacity_rank_and_id() {
.unwrap();
assert_eq!(
resolver
.resolve(10, None)
.resolve(10, None, None, 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");
assert_eq!(resolver.resolve(11, None, None, None).unwrap()[0].id, "min");
assert_eq!(resolver.resolve(15, None, None, None).unwrap()[0].id, "min");
assert_eq!(resolver.resolve(20, None, None, None).unwrap()[0].id, "a");
assert_eq!(
resolver.resolve(21, None, None, None).unwrap()[0].id,
"catch-all"
);
}
#[test]
@@ -207,17 +210,29 @@ fn context_capacity_checks_peak_when_known_and_input_otherwise() {
let mut long = bucket("long", None, policy);
long.max_context_tokens = Some(30);
let resolver = BucketResolver::new(vec![long, short]).unwrap();
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_eq!(
resolver.resolve(10, None, None, None).unwrap()[0].id,
"short"
);
assert_eq!(
resolver.resolve(10, Some(20), None, None).unwrap()[0].id,
"short"
);
assert_eq!(
resolver.resolve(10, Some(21), None, None).unwrap()[0].id,
"long"
);
assert!(resolver
.resolve(10, Some(31), None, None)
.unwrap()
.is_empty());
assert!(resolver.resolve(31, None, None, None).unwrap().is_empty());
assert!(matches!(
resolver.resolve(10, Some(9)),
resolver.resolve(10, Some(9), None, None),
Err(PickError::InvalidSignal(_))
));
assert!(BucketResolver::default()
.resolve(1, None)
.resolve(1, None, None, None)
.unwrap()
.is_empty());
}
@@ -238,7 +253,7 @@ async fn selected_pd_bucket_owns_both_memberships_and_policies() {
},
)])
.unwrap();
let bucket = resolver.resolve(10, Some(20)).unwrap()[0];
let bucket = resolver.resolve(10, Some(20), None, None).unwrap()[0];
let request = BucketRequest {
prefix: None,
model: &model,
@@ -270,7 +285,7 @@ async fn resolver_includes_empty_groups_without_invoking_policies() {
};
let resolver =
BucketResolver::new(vec![empty, bucket("available", Some(20), policy.clone())]).unwrap();
let buckets = resolver.resolve(10, None).unwrap();
let buckets = resolver.resolve(10, None, None, None).unwrap();
assert_eq!(
buckets
.iter()
@@ -0,0 +1,135 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use std::sync::Arc;
use sgl_router::buckets_reorg::{Bucket, BucketGroups, BucketResolver, EngineGroup, SloPreference};
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
use sgl_router::policies_reorg::PickError;
use sgl_router::state::load_monitor::engine_reported_load::EngineReportedLoadTable;
fn bucket(id: &str, max: u64, rank: u32, ttft: Option<u64>, tps: Option<f64>) -> Bucket {
let mut bucket = Bucket::new(
id,
BucketGroups::Plain(EngineGroup::new(Arc::new(PowerOfTwoPolicy::new(
EngineReportedLoadTable::new(),
)))),
);
bucket.limits.max = Some(max);
bucket.rank = rank;
bucket.ttft_ms = ttft;
bucket.tokens_per_second = tps;
bucket
}
fn resolver() -> BucketResolver {
BucketResolver::new(vec![
bucket("both", 100, 9, Some(50), Some(100.0)),
bucket("ttft", 10, 1, Some(50), Some(10.0)),
bucket("tps", 10, 2, Some(100), Some(100.0)),
bucket("neither", 8, 0, Some(100), Some(10.0)),
bucket("missing", 10, 3, None, None),
bucket("too-short", 1, 0, Some(1), Some(1000.0)),
])
.unwrap()
}
fn ids(resolver: &BucketResolver, ttft: Option<u64>, tps: Option<f64>) -> Vec<&str> {
resolver
.resolve(5, Some(5), ttft, tps)
.unwrap()
.into_iter()
.map(|b| b.id.as_str())
.collect()
}
#[test]
fn slo_tiers_keep_length_constraints_and_capacity_rank_order() {
let mut resolver = resolver();
assert_eq!(
ids(&resolver, Some(50), Some(100.0)),
["neither", "ttft", "tps", "missing", "both"]
);
resolver.ttft_slo = SloPreference::SloFirst;
resolver.tps_slo = SloPreference::SloFirst;
assert_eq!(
ids(&resolver, Some(50), Some(100.0)),
["both", "ttft", "tps", "neither", "missing"]
);
resolver.ttft_slo = SloPreference::BestEffort;
resolver.tps_slo = SloPreference::BestEffort;
assert_eq!(
ids(&resolver, Some(50), Some(100.0)),
["neither", "missing", "ttft", "tps", "both"]
);
resolver.ttft_slo = SloPreference::SloFirst;
assert_eq!(
ids(&resolver, Some(50), Some(100.0)),
["ttft", "neither", "missing", "both", "tps"]
);
}
#[test]
fn absent_targets_are_neutral_and_each_preference_can_be_disabled() {
let mut resolver = resolver();
resolver.ttft_slo = SloPreference::SloFirst;
resolver.tps_slo = SloPreference::BestEffort;
assert_eq!(
ids(&resolver, None, None),
["neither", "ttft", "tps", "missing", "both"]
);
assert_eq!(
ids(&resolver, Some(50), None),
["ttft", "both", "neither", "tps", "missing"]
);
resolver.ttft_slo = SloPreference::Disabled;
resolver.tps_slo = SloPreference::SloFirst;
assert_eq!(
ids(&resolver, Some(1), Some(100.0)),
["tps", "both", "neither", "ttft", "missing"]
);
}
#[test]
fn invalid_enabled_targets_fail_and_invalid_estimates_do_not_match() {
let mut resolver = resolver();
resolver.ttft_slo = SloPreference::SloFirst;
resolver.tps_slo = SloPreference::SloFirst;
assert!(matches!(
resolver.resolve(5, None, Some(0), None),
Err(PickError::InvalidSignal(_))
));
for tps in [0.0, -1.0, f64::NAN, f64::INFINITY] {
assert!(matches!(
resolver.resolve(5, None, None, Some(tps)),
Err(PickError::InvalidSignal(_))
));
}
resolver.buckets[0].ttft_ms = Some(0);
resolver.buckets[0].tokens_per_second = Some(f64::INFINITY);
assert_eq!(
ids(&resolver, Some(50), Some(100.0)),
["ttft", "tps", "neither", "missing", "both"]
);
resolver.ttft_slo = SloPreference::Disabled;
resolver.tps_slo = SloPreference::Disabled;
assert!(resolver.resolve(5, None, Some(0), Some(f64::NAN)).is_ok());
}
#[test]
fn slo_preferences_never_relax_peak_capacity_or_input_range() {
let mut resolver = resolver();
resolver.ttft_slo = SloPreference::SloFirst;
resolver.tps_slo = SloPreference::SloFirst;
resolver.buckets[0].max_context_tokens = Some(6);
resolver.buckets[1].limits.min = Some(6);
let found = resolver.resolve(5, Some(7), Some(50), Some(100.0)).unwrap();
assert_eq!(
found.iter().map(|b| b.id.as_str()).collect::<Vec<_>>(),
["tps", "neither", "missing"]
);
assert!(matches!(
resolver.resolve(5, Some(4), None, None),
Err(PickError::InvalidSignal(_))
));
}
@@ -14,6 +14,7 @@ use sgl_router::server::app_context::ChatRouting;
use std::sync::Mutex;
mod session_aware;
mod slo;
type PickCall = (String, Stage, u64, Option<u64>);
@@ -0,0 +1,134 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use super::*;
#[tokio::test]
async fn slo_headers_order_complete_pd_buckets_and_validate_before_dispatch() {
use sgl_router::buckets_reorg::SloPreference;
let fast_p = MockWorker::start(vec![]).await;
let slow_p = MockWorker::start(vec![]).await;
let fast_d = MockWorker::start(vec![]).await;
let slow_d = MockWorker::start(vec![]).await;
let mut ctx = Arc::try_unwrap(context(
&[
("fast-p", Stage::Prefill, &fast_p),
("slow-p", Stage::Prefill, &slow_p),
("fast-d", Stage::Decode, &fast_d),
("slow-d", Stage::Decode, &slow_d),
],
vec![],
))
.unwrap_or_else(|_| panic!("context is not shared yet"));
let mut prefill = Bucket::new(
"a-prefill",
BucketGroups::Pd {
prefill: group("fast-p", Arc::new(FirstPolicy::default())),
decode: group("slow-d", Arc::new(FirstPolicy::default())),
},
);
prefill.ttft_ms = Some(50);
prefill.tokens_per_second = Some(10.0);
let mut decode = Bucket::new(
"b-decode",
BucketGroups::Pd {
prefill: group("slow-p", Arc::new(FirstPolicy::default())),
decode: group("fast-d", Arc::new(FirstPolicy::default())),
},
);
decode.ttft_ms = Some(100);
decode.tokens_per_second = Some(100.0);
let mut resolver = BucketResolver::new(vec![decode, prefill]).unwrap();
resolver.ttft_slo = SloPreference::SloFirst;
resolver.tps_slo = SloPreference::SloFirst;
ctx.chat_routing = ChatRouting::Reorg([(ModelId("tiny".into()), resolver)].into());
let app = build_router(Arc::new(ctx));
for (ttft, tps) in [("0", "100"), ("50", "NaN"), ("50", "0"), ("abc", "100")] {
let mut req = request(body("hi"));
req.headers_mut()
.insert("x-sgl-ttft-slo-ms", ttft.parse().unwrap());
req.headers_mut()
.insert("x-sgl-tps-slo", tps.parse().unwrap());
assert_eq!(
app.clone().oneshot(req).await.unwrap().status(),
StatusCode::BAD_REQUEST
);
}
for worker in [&fast_p, &slow_p, &fast_d, &slow_d] {
assert!(worker.captured.lock().unwrap().last_body.is_none());
}
// With both targets unmet once, the existing bucket order breaks the tie.
// With TTFT unconstrained, throughput chooses the other complete bucket.
for (ttft, expected_d, expected_p, unused_d, unused_p) in [
("50", &slow_d, &fast_p, &fast_d, &slow_p),
("100", &fast_d, &slow_p, &slow_d, &fast_p),
] {
for worker in [&fast_p, &slow_p, &fast_d, &slow_d] {
worker.captured.lock().unwrap().last_body = None;
}
let mut req = request(body("hi"));
req.headers_mut()
.insert("x-sgl-ttft-slo-ms", ttft.parse().unwrap());
req.headers_mut()
.insert("x-sgl-tps-slo", "100".parse().unwrap());
let response = app.clone().oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()["x-sgl-decode-url"], expected_d.url);
response.into_body().collect().await.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
while expected_p.captured.lock().unwrap().last_body.is_none() {
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.unwrap();
assert!(unused_p.captured.lock().unwrap().last_body.is_none());
assert!(unused_d.captured.lock().unwrap().last_body.is_none());
}
}
#[tokio::test]
async fn preferred_bucket_rejection_falls_back_and_disabled_headers_are_ignored() {
use sgl_router::buckets_reorg::SloPreference;
for enabled in [false, true] {
let rejected = MockWorker::start(vec![]).await;
let accepted = MockWorker::start(vec![]).await;
let mut ctx = Arc::try_unwrap(context(
&[
("rejected", Stage::Plain, &rejected),
("accepted", Stage::Plain, &accepted),
],
vec![],
))
.unwrap_or_else(|_| panic!("context is not shared yet"));
let mut fast = Bucket::new(
"a-fast",
BucketGroups::Plain(group("rejected", rejecting_policy())),
);
fast.ttft_ms = Some(10);
let mut slow = Bucket::new(
"b-slow",
BucketGroups::Plain(group("accepted", Arc::new(FirstPolicy::default()))),
);
slow.ttft_ms = Some(100);
let mut resolver = BucketResolver::new(vec![slow, fast]).unwrap();
if enabled {
resolver.ttft_slo = SloPreference::SloFirst;
}
ctx.chat_routing = ChatRouting::Reorg([(ModelId("tiny".into()), resolver)].into());
let app = build_router(Arc::new(ctx));
let mut req = request(body("hi"));
req.headers_mut().insert(
"x-sgl-ttft-slo-ms",
if enabled { "50" } else { "invalid" }.parse().unwrap(),
);
// TPS is disabled even when TTFT is enabled.
req.headers_mut()
.insert("x-sgl-tps-slo", "NaN".parse().unwrap());
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
response.into_body().collect().await.unwrap();
assert!(rejected.captured.lock().unwrap().last_body.is_none());
assert!(accepted.captured.lock().unwrap().last_body.is_some());
}
}