[sgl-router] refactor - generalized admission policy definitions (#40271)
Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5.1
parent
70b5b03e78
commit
a9871012ac
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user