Files
sglang/experimental/sgl-router/tests/proxy/pd_pool_isolation.rs
T

526 lines
20 KiB
Rust

// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! PD pool isolation — end-to-end at the HTTP layer using MockWorker.
//!
//! Drives the chat handler with:
//!
//! * A model whose registered workers are all `WorkerMode::Decode`. The
//! handler dispatches **prefill** traffic (chat-completions is the
//! prefill phase of a PD request), so it must return 503 with
//! `no_prefill_workers_available`.
//! * A model with no workers at all → 503 `no_healthy_workers`
//! (existing code path; pinned here so a future PD wiring change
//! doesn't silently swap codes).
//! * A PD-disagg model with both pools healthy → request flows to the
//! prefill worker (sanity check; the decode worker MUST NOT be selected for
//! the chat route).
use axum::body::Body;
use axum::http::{Request, StatusCode};
use http_body_util::BodyExt;
use sgl_router::config::{
Config, DiscoveryBackend, InflightLoadConfig, ModelConfig, ObservabilityConfig, PolicyKind,
ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults;
use sgl_router::proxy::Proxy;
use sgl_router::server::app::build_router;
use sgl_router::server::app_context::AppContext;
use sgl_router::tokenizer::TokenizerRegistry;
use sgl_router::workers::WorkerRegistry;
use std::sync::Arc;
use std::time::Duration;
use tower::ServiceExt;
fn config() -> Config {
Config {
server: ServerConfig {
host: "0".into(),
port: 0,
..Default::default()
},
observability: ObservabilityConfig::default(),
model: ModelConfig {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
disable_input_ids_forwarding: false,
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
affinity: None,
fused: None,
eligibility: None,
sampling_overrides: Default::default(),
},
discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()],
}),
proxy: ProxyConfig::default(),
router_inflight_load: InflightLoadConfig::default(),
}
}
fn build_ctx(specs: Vec<WorkerSpec>) -> Arc<AppContext> {
let cfg = config();
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let registry = Arc::new(WorkerRegistry::default());
for s in specs {
let _ = registry.add(s);
}
let policies = Arc::new(build_registry_with_defaults(&cfg).unwrap());
let proxy = Arc::new(Proxy::new(Duration::from_secs(5)).unwrap());
Arc::new(AppContext::new(cfg, tokenizers, proxy, registry, policies))
}
fn chat_request() -> Request<Body> {
Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({
"model": "tiny",
"messages": [{"role": "user", "content": "hi"}],
}))
.unwrap(),
))
.unwrap()
}
#[tokio::test]
async fn pd_decode_stream_expires_after_prefill_completes() {
use sgl_router::state::load_monitor::router_inflight_load::{
MockClock, RouterInflightLoadRegistry,
};
use std::time::Instant;
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start_slow_stream(
vec!["data: chunk\n\n"; 1000],
Duration::from_millis(10),
)
.await;
let cfg = config();
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let registry = Arc::new(WorkerRegistry::default());
registry
.add(WorkerSpec {
id: WorkerId("p1".into()),
url: prefill.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: Some(8997),
})
.unwrap();
registry
.add(WorkerSpec {
id: WorkerId("d1".into()),
url: decode.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
})
.unwrap();
let prefill_worker = registry.get(&WorkerId("p1".into())).unwrap();
let decode_worker = registry.get(&WorkerId("d1".into())).unwrap();
let clock = Arc::new(MockClock::new(Instant::now()));
let inflight = RouterInflightLoadRegistry::new(clock.clone(), Duration::from_secs(10));
let policies = Arc::new(build_registry_with_defaults(&cfg).unwrap());
let proxy = Arc::new(Proxy::new(Duration::from_secs(5)).unwrap());
let ctx = Arc::new(AppContext::with_router_inflight_load(
cfg,
tokenizers,
proxy,
registry,
policies,
inflight.clone(),
));
let request = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(
serde_json::json!({
"model": "tiny",
"messages": [{"role": "user", "content": "hi"}],
"stream": true,
})
.to_string(),
))
.unwrap();
// Two prior faults make any accidental expiry failure trip the default breaker.
decode_worker.breaker.record_failure();
decode_worker.breaker.record_failure();
let response = build_router(ctx.clone()).oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let mut body = response.into_body();
tokio::time::timeout(Duration::from_secs(2), body.frame())
.await
.unwrap()
.unwrap()
.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
while inflight.inflight_count() != 1 || prefill_worker.router_inflight_load() != 0 {
tokio::task::yield_now().await;
}
})
.await
.expect("prefill must complete before expiring decode");
assert_eq!(decode_worker.router_inflight_load(), 1);
clock.advance(Duration::from_secs(11));
assert_eq!(inflight.sweep_stale(), 1);
let error = tokio::time::timeout(Duration::from_secs(2), body.collect())
.await
.expect("decode stream must stop when its registration expires")
.unwrap_err();
assert!(error.to_string().contains("stale_request_timeout"));
assert_eq!(decode_worker.router_inflight_load(), 0);
assert_eq!(decode_worker.breaker.snapshot().state_code, 0);
let metrics = ctx.metrics.render();
assert!(metrics
.lines()
.any(|line| line == r#"sgl_router_stale_requests_total{outcome="expired"} 1"#));
let expected = format!(
r#"sgl_router_stream_outcome_total{{worker_url="{}",model_id="tiny",outcome="expired"}} 1"#,
decode.url,
);
assert!(metrics.lines().any(|line| line == expected));
assert!(!metrics
.lines()
.any(|line| line.starts_with("sgl_router_stream_outcome_total{")
&& line.contains(r#"outcome="upstream_error""#)));
decode_worker.breaker.record_failure();
assert_eq!(decode_worker.breaker.snapshot().state_code, 1);
}
/// Gap closer #1: PD mode with only decode workers → 503 with
/// `no_prefill_workers_available`. The chat route is a prefill
/// dispatch, so a decode-only pool means partial failure.
#[tokio::test]
async fn pd_mode_decode_only_returns_no_prefill_workers_available() {
let worker = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![WorkerSpec {
id: WorkerId("d1".into()),
url: worker.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
}]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
res.headers().get("x-router-error-code").unwrap(),
"no_prefill_workers_available",
);
let body = res.into_body().collect().await.unwrap().to_bytes();
let body_str = String::from_utf8_lossy(&body);
assert!(
body_str.contains("\"code\":\"no_prefill_workers_available\""),
"body: {body_str}"
);
}
/// Pin the existing-code-path branch: no workers at all → 503 with
/// `no_healthy_workers`. Ensures the new PD code path didn't swap the
/// code for the "model has zero workers" case.
#[tokio::test]
async fn no_workers_returns_no_healthy_workers() {
let ctx = build_ctx(vec![]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
res.headers().get("x-router-error-code").unwrap(),
"no_healthy_workers",
);
}
/// PD-disagg deployment with both pools healthy → chat dispatch fans
/// out to BOTH the prefill and the decode worker (Pattern B: prefill
/// in a detached task, decode awaited for the client response). Both
/// receive the same bootstrap-injected body so the SGLang engine can
/// match KV transfers via `bootstrap_room`. Pool *isolation* — the
/// guarantee that the policy's prefill candidate set excludes decode
/// workers — is exercised at the resolver layer
/// (`policies::registry::tests::pd_resolution_returns_distinct_pools`).
/// Here we only assert the HTTP-layer wiring of the dual dispatch.
#[tokio::test]
async fn pd_mode_chat_dispatch_fans_to_both_prefill_and_decode() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![
WorkerSpec {
id: WorkerId("p1".into()),
url: prefill.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: Some(8997),
},
WorkerSpec {
id: WorkerId("d1".into()),
url: decode.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
]);
let app = build_router(ctx);
// Fire a single request; both prefill (spawn-and-forget) and
// decode (awaited) must receive a body with the injected
// bootstrap fields. The decode body is what the client sees on
// the response.
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"decode response status should reach the client",
);
// Decode receives its body synchronously (we awaited it), so it's
// guaranteed captured by the time the response returned. Scope
// the lock guard to this block so it doesn't span the `.await`
// below (clippy: await_holding_lock).
{
let decode_seen = decode.captured.lock().unwrap();
assert!(
decode_seen.last_body.is_some(),
"decode worker must receive the bootstrap-injected request body in PD mode",
);
}
// Prefill is detached; poll briefly until its capture lands. The
// prefill task races the HTTP response back to the client. The
// local binding releases the `std::sync::Mutex` guard before the
// `.await` — holding a sync mutex across an await would let one
// task pin the lock while another tries to acquire it.
let prefill_body = tokio::time::timeout(Duration::from_secs(2), async {
loop {
let captured = prefill.captured.lock().unwrap().last_body.clone();
if let Some(b) = captured {
return b;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.expect("prefill MUST eventually receive its body via the detached task");
assert!(!prefill_body.is_empty());
}
/// PD-mode chat request carries an `x-sgl-decode-url` header for the final
/// Decode decision. Step 1 defaults to Decode P2; the header remains an
/// observability contract regardless of which Decode policy produced it.
#[tokio::test]
async fn pd_mode_chat_dispatch_sets_final_decode_header() {
use std::collections::HashSet;
let prefill_a = crate::common::mock_worker::MockWorker::start(vec![]).await;
let prefill_b = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode_a = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode_b = crate::common::mock_worker::MockWorker::start(vec![]).await;
// MockWorker URLs all bind to `127.0.0.1`; this test deliberately does
// not assert a host relation. It pins only the HTTP wiring: the final D
// selected by the role-local policy is reflected on the P request.
let ctx = build_ctx(vec![
WorkerSpec {
id: WorkerId("p1".into()),
url: prefill_a.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
WorkerSpec {
id: WorkerId("p2".into()),
url: prefill_b.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
WorkerSpec {
id: WorkerId("d1".into()),
url: decode_a.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
WorkerSpec {
id: WorkerId("d2".into()),
url: decode_b.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
]);
let app = build_router(ctx);
// Fire 4 requests; both prefill workers see traffic via round-robin.
for _ in 0..4 {
let res = app.clone().oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
// Every request that hit a prefill mock MUST carry the final-decode
// header. The value MUST be one of the two registered Decode URLs.
let decode_urls: HashSet<String> = [decode_a.url.clone(), decode_b.url.clone()]
.into_iter()
.collect();
for (label, p) in [("prefill_a", &prefill_a), ("prefill_b", &prefill_b)] {
let g = p.captured.lock().unwrap();
if g.last_body.is_none() {
// This prefill didn't receive a request — round-robin's
// dashmap iteration is non-deterministic, so one side may
// skip in a 4-request fire. Continue.
continue;
}
let hdr = g.headers.get("x-sgl-decode-url").unwrap_or_else(|| {
panic!(
"{label} did not receive an x-sgl-decode-url header. headers: {:?}",
g.headers
)
});
assert!(
decode_urls.contains(hdr),
"{label} got decode hint {hdr}, expected one of {decode_urls:?}",
);
}
}
/// Task C: plain-mode (non-PD) request does NOT carry the
/// `x-sgl-decode-url` header. Pin: the affinity step is gated on
/// `worker.mode() == Prefill` so plain workers are not asked to
/// bootstrap nonexistent decode peers.
#[tokio::test]
async fn plain_mode_chat_dispatch_omits_decode_affinity_header() {
let plain = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![WorkerSpec {
id: WorkerId("w1".into()),
url: plain.url.clone(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
}]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let g = plain.captured.lock().unwrap();
assert!(
!g.headers.contains_key("x-sgl-decode-url"),
"plain-mode worker must not receive a decode-affinity header. headers: {:?}",
g.headers,
);
}
/// Task C: PD-mode prefill request with NO decode workers → 503
/// `no_decode_workers_available`. Pin: failure mode is loud and
/// distinct from the existing `no_prefill_workers_available` path.
#[tokio::test]
async fn pd_mode_prefill_only_returns_no_decode_workers_available() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![WorkerSpec {
id: WorkerId("p1".into()),
url: prefill.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
}]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
res.headers().get("x-router-error-code").unwrap(),
"no_decode_workers_available",
);
}
/// PD-mode chat response carries `x-sgl-decode-url` so external tests
/// can observe final Decode selection end-to-end (without sniffing the proxy
/// hop into the upstream prefill worker). Mirrors the request-side
/// behavior asserted by `pd_mode_chat_dispatch_sets_final_decode_header`.
#[tokio::test]
async fn pd_mode_chat_response_carries_decode_affinity_header() {
use std::collections::HashSet;
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode_a = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode_b = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![
WorkerSpec {
id: WorkerId("p1".into()),
url: prefill.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
WorkerSpec {
id: WorkerId("d1".into()),
url: decode_a.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
WorkerSpec {
id: WorkerId("d2".into()),
url: decode_b.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let decode_urls: HashSet<String> = [decode_a.url.clone(), decode_b.url.clone()]
.into_iter()
.collect();
let hdr = res
.headers()
.get("x-sgl-decode-url")
.unwrap_or_else(|| {
panic!(
"PD-mode chat response did not carry x-sgl-decode-url; headers: {:?}",
res.headers(),
)
})
.to_str()
.unwrap()
.to_owned();
assert!(
decode_urls.contains(&hdr),
"response carried decode hint {hdr}, expected one of {decode_urls:?}",
);
}
/// Plain-mode chat response does NOT carry `x-sgl-decode-url`. Pin: the
/// response-side mirror is gated on PD-mode dispatch.
#[tokio::test]
async fn plain_mode_chat_response_omits_decode_affinity_header() {
let plain = crate::common::mock_worker::MockWorker::start(vec![]).await;
let ctx = build_ctx(vec![WorkerSpec {
id: WorkerId("w1".into()),
url: plain.url.clone(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
}]);
let app = build_router(ctx);
let res = app.oneshot(chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert!(
!res.headers().contains_key("x-sgl-decode-url"),
"plain-mode chat response must not carry x-sgl-decode-url; headers: {:?}",
res.headers(),
);
}