[sgl-router] refactor - generalized admission policy definitions (#40271)

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
Kan Wu
2026-09-21 16:26:30 +08:00
committed by GitHub
co-authored by Claude Fable 5.1
parent 70b5b03e78
commit a9871012ac
12 changed files with 445 additions and 172 deletions
@@ -11,6 +11,7 @@ mod discovery;
mod health;
mod policies;
mod policies_reorg;
mod policies_reorg_admission;
mod policies_reorg_cache_aware;
mod policies_reorg_load;
mod policies_reorg_power_of_two;
@@ -8,9 +8,10 @@ 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::admission::{
AdmissionLimits, Decision, EngineAdmission, EngineMetrics,
};
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)]
@@ -25,7 +26,7 @@ struct TestPolicy {
impl Default for TestPolicy {
fn default() -> Self {
Self {
admission: Arc::new(AllowAll),
admission: Arc::new(AdmissionLimits::default()),
result: None,
miss: false,
invalid: false,
@@ -52,7 +53,9 @@ impl Policy for TestPolicy {
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)? {
if let Decision::Reject(reason) =
self.admission.check(&engine, &EngineMetrics::default())?
{
return Err(PickError::AdmissionRejected(Rejection {
engine: engine.id.clone(),
reason,
@@ -70,12 +73,7 @@ impl Policy for TestPolicy {
struct Reject(&'static str);
impl EngineAdmission for Reject {
fn check(
&self,
engine: &Worker,
_: &PickRequest<'_>,
_: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, _: &EngineMetrics) -> Result<Decision, PickError> {
Ok(if engine.id.0 == self.0 {
Decision::Reject("full".into())
} else {
@@ -437,12 +435,7 @@ async fn power_of_two_checks_selected_engine_and_propagates_rejection_without_fa
}
impl EngineAdmission for Check {
fn check(
&self,
engine: &Worker,
_: &PickRequest<'_>,
_: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, _: &EngineMetrics) -> Result<Decision, PickError> {
self.calls.lock().unwrap().push(engine.id.clone());
if self.invalid {
Err(PickError::InvalidSignal("admission input".into()))
@@ -0,0 +1,180 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use std::sync::Arc;
use std::time::Instant;
use serde_json::json;
use sgl_router::discovery::{ModelId, WorkerId, WorkerSpec};
use sgl_router::policies_reorg::admission::{
AdmissionLimits, Decision, EngineAdmission, EngineMetrics,
};
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, LoadStat, NativeCacheRankLoad,
};
use sgl_router::workers::Worker;
fn engine() -> Arc<Worker> {
Arc::new(Worker::new(WorkerSpec {
id: WorkerId("w".into()),
url: "http://w".into(),
mode: Stage::Plain,
model_ids: vec![ModelId("m".into())],
bootstrap_port: None,
}))
}
fn report(running: u64, waiting: u64, kv_tokens: u64, pending: u64) -> LoadStat {
LoadStat {
num_running_reqs: running,
num_waiting_reqs: waiting,
num_tokens: kv_tokens,
max_total_num_tokens: 1000,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: pending,
num_total_tokens: kv_tokens,
max_running_requests: 100,
total_prefill_uncached_tokens: 0,
total_prefill_busy_us: 0,
}),
}
}
fn reject(name: &str) -> Decision {
Decision::Reject(name.into())
}
#[test]
fn limits_are_flat_optional_fields() {
let limits: AdmissionLimits =
serde_json::from_value(json!({"max_inflight_requests": 4, "max_kv_tokens": 100})).unwrap();
assert_eq!(
limits,
AdmissionLimits {
max_inflight_requests: Some(4),
max_kv_tokens: Some(100),
..Default::default()
}
);
assert_eq!(
serde_json::from_value::<AdmissionLimits>(json!({})).unwrap(),
AdmissionLimits::default()
);
assert!(serde_json::from_value::<AdmissionLimits>(json!({"max_in_flight": 4})).is_err());
}
#[test]
fn each_limit_caps_its_metric_and_fails_open_when_unknown() {
let engine = engine();
let known = EngineMetrics {
running_requests: Some(3),
waiting_requests: Some(1),
kv_tokens: Some(80),
pending_prefill_tokens: Some(90),
inflight_requests: 2,
};
let unknown = EngineMetrics {
inflight_requests: 2,
..Default::default()
};
let limit = |field: &str, max| {
let mut limits = AdmissionLimits::default();
*match field {
"max_running_requests" => &mut limits.max_running_requests,
"max_waiting_requests" => &mut limits.max_waiting_requests,
"max_kv_tokens" => &mut limits.max_kv_tokens,
"max_pending_prefill_tokens" => &mut limits.max_pending_prefill_tokens,
"max_inflight_requests" => &mut limits.max_inflight_requests,
_ => unreachable!(),
} = Some(max);
limits
};
assert_eq!(
AdmissionLimits::default().check(&engine, &known).unwrap(),
Decision::Allow
);
for (field, below, at) in [
("max_running_requests", 4, 3),
("max_waiting_requests", 2, 1),
("max_kv_tokens", 81, 80),
("max_pending_prefill_tokens", 91, 90),
("max_inflight_requests", 3, 2),
] {
let (below, at) = (limit(field, below), limit(field, at));
assert_eq!(below.check(&engine, &known).unwrap(), Decision::Allow);
assert_eq!(at.check(&engine, &known).unwrap(), reject(field));
// Without a report only the router-local in-flight count applies.
let expected = if field == "max_inflight_requests" {
reject(field)
} else {
Decision::Allow
};
assert_eq!(at.check(&engine, &unknown).unwrap(), expected);
}
}
#[tokio::test]
async fn metrics_come_from_the_selection_snapshot_and_live_inflight_count() {
let engine = engine();
let table = EngineReportedLoadTable::new();
table.set(&engine.url, 0, report(1, 2, 80, 5), Instant::now());
let guard = engine.load_guard();
assert_eq!(
EngineMetrics::observe(&engine, &table.capture_snapshot(Instant::now())),
EngineMetrics {
running_requests: Some(1),
waiting_requests: Some(2),
kv_tokens: Some(80),
pending_prefill_tokens: Some(5),
inflight_requests: 1,
}
);
drop(guard);
// Basic reports carry request counts but no KV or pending-prefill tokens.
let basic = LoadStat {
native_cache: None,
..report(1, 2, 80, 5)
};
table.set(&engine.url, 0, basic, Instant::now());
let metrics = EngineMetrics::observe(&engine, &table.capture_snapshot(Instant::now()));
assert_eq!(
(
metrics.running_requests,
metrics.kv_tokens,
metrics.pending_prefill_tokens
),
(Some(1), None, None)
);
let mut policy = PowerOfTwoPolicy::new(table.clone());
policy.admission = Arc::new(AdmissionLimits {
max_running_requests: Some(2),
max_kv_tokens: Some(100),
..Default::default()
});
let model = ModelId("m".into());
let request = PickRequest::new(&model, Stage::Plain, 10);
let engines = [engine];
for (running, kv_tokens, rejection) in [
(1, 100, Some("max_kv_tokens")),
(1, 99, None),
(2, 0, Some("max_running_requests")),
] {
table.set(
&engines[0].url,
0,
report(running, 0, kv_tokens, 0),
Instant::now(),
);
let result = policy.pick(&engines, &request).await;
match rejection {
None => assert!(result.is_ok()),
Some(reason) => assert!(matches!(
result,
Err(PickError::AdmissionRejected(rejected)) if rejected.reason == reason
)),
}
}
}
@@ -10,7 +10,7 @@ use sgl_router::buckets_reorg::{Bucket, BucketGroups, BucketResolver, EngineGrou
use sgl_router::config::AffinityConfig;
use sgl_router::discovery::{ModelId, WorkerId, WorkerSpec};
use sgl_router::policies::prefix_provider::RadixTreePrefixProvider;
use sgl_router::policies_reorg::admission::{Decision, EngineAdmission};
use sgl_router::policies_reorg::admission::{Decision, EngineAdmission, EngineMetrics};
use sgl_router::policies_reorg::cache_aware::{CacheAwarePolicy, CacheSource, PrefixMemo};
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
use sgl_router::policies_reorg::{PickError, PickRequest, Policy, Stage};
@@ -18,7 +18,7 @@ use sgl_router::state::kv_events::{
compute_block_hashes, compute_block_hashes_bigram, BlockSizeOracle, HashTree, KvWorkerId,
};
use sgl_router::state::load_monitor::engine_reported_load::{
EngineReportedLoadTable, EngineReportedWorkerLoad, LoadStat, NativeCacheRankLoad,
EngineReportedLoadTable, LoadStat, NativeCacheRankLoad,
};
use sgl_router::workers::Worker;
@@ -116,16 +116,11 @@ impl Reject {
}
impl EngineAdmission for Reject {
fn check(
&self,
engine: &Worker,
_: &PickRequest<'_>,
load: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, metrics: &EngineMetrics) -> Result<Decision, PickError> {
self.calls
.lock()
.unwrap()
.push((engine.id.0.clone(), load.map(|load| load.num_waiting_reqs)));
.push((engine.id.0.clone(), metrics.waiting_requests));
Ok(if engine.id.0 == self.id {
Decision::Reject("full".into())
} else {
@@ -5,12 +5,10 @@ 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::admission::{Decision, EngineAdmission, EngineMetrics};
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::state::load_monitor::engine_reported_load::{EngineReportedLoadTable, LoadStat};
use sgl_router::workers::Worker;
const URL: &str = "http://engine";
@@ -43,21 +41,16 @@ fn engine() -> Arc<Worker> {
#[derive(Debug)]
struct ObserveAdmission {
table: Arc<EngineReportedLoadTable>,
observations: Mutex<Vec<Option<EngineReportedWorkerLoad>>>,
observations: Mutex<Vec<EngineMetrics>>,
}
impl EngineAdmission for ObserveAdmission {
fn check(
&self,
engine: &Worker,
_: &PickRequest<'_>,
load: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, metrics: &EngineMetrics) -> 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());
self.observations.lock().unwrap().push(*metrics);
Ok(Decision::Allow)
}
}
@@ -95,17 +88,16 @@ async fn selected_load_reaches_admission_and_next_pick_reads_fresh_state() {
}
let observations = admission.observations.lock().unwrap();
assert_eq!(observations.len(), 2);
// Basic reports carry request counts but no KV or pending-prefill tokens.
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,
})
EngineMetrics {
running_requests: Some(4),
waiting_requests: Some(4),
..EngineMetrics::default()
}
);
assert_eq!(observations[1].as_ref().unwrap().num_running_reqs, 102);
assert_eq!(observations[1].running_requests, Some(102));
}
#[tokio::test]
@@ -133,6 +125,10 @@ async fn missing_stale_and_incomplete_reports_reach_admission_as_unknown() {
.pick(&[engine()], &PickRequest::new(&model, Stage::Plain, 10))
.await
.unwrap();
assert_eq!(*admission.observations.lock().unwrap(), [None], "{case}");
assert_eq!(
*admission.observations.lock().unwrap(),
[EngineMetrics::default()],
"{case}"
);
}
}
@@ -6,12 +6,10 @@ 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::admission::{Decision, EngineAdmission, EngineMetrics};
use sgl_router::policies_reorg::session_aware::SessionAwarePolicy;
use sgl_router::policies_reorg::{PickError, PickRequest, Policy, Stage};
use sgl_router::state::load_monitor::engine_reported_load::{
EngineReportedLoadTable, EngineReportedWorkerLoad, LoadStat,
};
use sgl_router::state::load_monitor::engine_reported_load::{EngineReportedLoadTable, LoadStat};
use sgl_router::state::load_monitor::router_inflight_load::MockClock;
use sgl_router::state::AffinityStore;
use sgl_router::workers::Worker;
@@ -52,16 +50,11 @@ struct Admission {
}
impl EngineAdmission for Admission {
fn check(
&self,
engine: &Worker,
_: &PickRequest<'_>,
load: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, metrics: &EngineMetrics) -> Result<Decision, PickError> {
self.calls
.lock()
.unwrap()
.push((engine.id.0.clone(), load.map(|load| load.num_waiting_reqs)));
.push((engine.id.0.clone(), metrics.waiting_requests));
if self.invalid.load(Ordering::Relaxed) {
return Err(PickError::InvalidSignal("invalid admission input".into()));
}
@@ -339,6 +332,7 @@ async fn shared_store_refreshes_active_sessions_and_expires_idle_ones() {
#[derive(Debug)]
struct RacingAdmission {
model: ModelId,
competitor: SessionAwarePolicy,
winner: Arc<Worker>,
reject_winner: bool,
@@ -346,19 +340,14 @@ struct RacingAdmission {
}
impl EngineAdmission for RacingAdmission {
fn check(
&self,
engine: &Worker,
request: &PickRequest<'_>,
_: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, engine: &Worker, _: &EngineMetrics) -> Result<Decision, PickError> {
self.calls.lock().unwrap().push(engine.id.0.clone());
if engine.id.0 == "a" {
// Complete a competing first request after this request selected a,
// but before it can commit. The effective binding is now b.
futures::executor::block_on(
self.competitor
.pick(std::slice::from_ref(&self.winner), request),
.pick(std::slice::from_ref(&self.winner), &request(&self.model)),
)?;
}
Ok(if self.reject_winner && engine.id == self.winner.id {
@@ -375,6 +364,7 @@ async fn concurrent_assignment_winner_is_checked_and_preserved_on_rejection() {
let (mut policy, store) = policy();
let engines = [engine("a", 0), engine("b", 9)];
let admission = Arc::new(RacingAdmission {
model: ModelId("m".into()),
competitor: SessionAwarePolicy::new(store.clone(), EngineReportedLoadTable::new()),
winner: engines[1].clone(),
reject_winner,
@@ -6,10 +6,11 @@ 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::admission::{
AdmissionLimits, Decision, EngineAdmission, EngineMetrics,
};
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;
mod session_aware;
@@ -27,7 +28,7 @@ struct FirstPolicy {
impl Default for FirstPolicy {
fn default() -> Self {
Self {
admission: Arc::new(AllowAll),
admission: Arc::new(AdmissionLimits::default()),
miss: false,
invalid: false,
calls: Mutex::new(Vec::new()),
@@ -58,7 +59,9 @@ impl Policy for FirstPolicy {
return Err(PickError::NoCandidates);
}
let engine = engines[0].clone();
if let Decision::Reject(reason) = self.admission.check(&engine, request, None)? {
if let Decision::Reject(reason) =
self.admission.check(&engine, &EngineMetrics::default())?
{
return Err(PickError::AdmissionRejected(Rejection {
engine: engine.id.clone(),
reason,
@@ -76,12 +79,7 @@ impl Policy for FirstPolicy {
struct RejectAll;
impl EngineAdmission for RejectAll {
fn check(
&self,
_: &Worker,
_: &PickRequest<'_>,
_: Option<&EngineReportedWorkerLoad>,
) -> Result<Decision, PickError> {
fn check(&self, _: &Worker, _: &EngineMetrics) -> Result<Decision, PickError> {
Ok(Decision::Reject("full".into()))
}
}
@@ -147,6 +145,70 @@ fn body(content: &str) -> serde_json::Value {
serde_json::json!({"model": "tiny", "messages": [{"role": "user", "content": content}]})
}
#[tokio::test]
async fn configured_limits_reject_before_dispatch_and_admit_after_load_drops() {
use sgl_router::policies_reorg::power_of_two::PowerOfTwoPolicy;
use sgl_router::state::load_monitor::engine_reported_load::{LoadStat, NativeCacheRankLoad};
use std::time::Instant;
let worker = MockWorker::start(vec![]).await;
let mut ctx = Arc::try_unwrap(context(&[("w", Stage::Plain, &worker)], vec![]))
.unwrap_or_else(|_| panic!("context is not shared yet"));
let mut policy = PowerOfTwoPolicy::new(ctx.engine_reported_load.clone());
policy.admission = Arc::new(AdmissionLimits {
max_running_requests: Some(1),
max_kv_tokens: Some(100),
..Default::default()
});
ctx.chat_routing = ChatRouting::Reorg(
[(
ModelId("tiny".into()),
BucketResolver::new(vec![Bucket::new(
"default",
BucketGroups::Plain(EngineGroup {
worker_ids: Some([WorkerId("w".into())].into()),
policy: Arc::new(policy),
}),
)]),
)]
.into(),
);
let ctx = Arc::new(ctx);
let app = build_router(ctx.clone());
for (running, kv_tokens, expected) in [
(1, 0, StatusCode::SERVICE_UNAVAILABLE),
(0, 100, StatusCode::SERVICE_UNAVAILABLE),
(0, 0, StatusCode::OK),
] {
ctx.engine_reported_load.set(
&worker.url,
0,
LoadStat {
num_running_reqs: running,
num_waiting_reqs: 0,
num_tokens: 0,
max_total_num_tokens: 1000,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 0,
num_total_tokens: kv_tokens,
max_running_requests: 10,
total_prefill_uncached_tokens: 0,
total_prefill_busy_us: 0,
}),
},
Instant::now(),
);
let response = app.clone().oneshot(request(body("hi"))).await.unwrap();
assert_eq!(response.status(), expected);
response.into_body().collect().await.unwrap();
assert_eq!(
worker.captured.lock().unwrap().last_body.is_some(),
expected == StatusCode::OK
);
assert_eq!(ctx.router_inflight_load.inflight_count(), 0);
}
}
#[tokio::test]
async fn length_selects_plain_bucket_before_engine_selection() {
let short_worker = MockWorker::start(vec![]).await;