[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
@@ -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;