[Router] Add bucket-aware policy domains and native cache indexing (#38108)

Signed-off-by: Vincent Gao <vincentbo@linux.alibaba.com>
Co-authored-by: inkcherry <mingzhi.liu@amd.com>
Co-authored-by: yangbodong22011 <13137470+yangbodong22011@users.noreply.github.com>
This commit is contained in:
Vincent Gao
2026-09-06 19:47:51 +08:00
committed by GitHub
co-authored by inkcherry yangbodong22011
parent a176ba2f7b
commit 5bebe7a033
76 changed files with 5842 additions and 3639 deletions
+16 -7
View File
@@ -57,18 +57,27 @@ sgl-router \
--model-id qwen3 \
--tokenizer-path /models/qwen3/tokenizer.json \
--worker-urls http://10.0.0.1:30000 http://10.0.0.2:30000 \
--policy cache_aware_zmq \
--policy cache_aware \
--cache-prefix-provider indexer \
--kv-indexer-endpoint http://10.0.0.10:50051 \
--kv-indexer-query-timeout-ms 100 \
--kv-indexer-query-max-inflight 32
```
The existing cache-aware policy and thresholds are reused. When configured, the
Indexer replaces the Router-local radix tree as the cache signal. A successful
query with no usable match selects by minimum active load; connection failures,
timeouts, local admission rejection, and server rejection fail the Router
request with `503` rather than silently switching signals. The timeout and local
concurrency bound default to 100ms and 32 respectively.
The Indexer replaces the Router-local radix tree as the native Cache-Aware
signal. Query timeouts and local concurrency are bounded by the two Indexer
options, which default to 100 ms and 32 respectively.
## Upgrading from `cache_aware_zmq`
The `cache_aware_zmq` policy has been removed. Configurations using it should
select `--policy cache_aware` and choose a native cache-prefix source: the
Router-local radix tree (the default), or the external Indexer shown above.
The legacy `--cache-threshold`, `--balance-abs-threshold`, and
`--balance-rel-threshold` flags have also been removed. They do not have
one-to-one replacements; remove them and review the current `sgl-router
--help` output when tuning Cache-Aware routing.
## License
@@ -5,10 +5,7 @@
//!
//! Mirrors `sgl-model-gateway/benches/manual_policy_benchmark.rs` —
//! measures how fast the routing layer returns a worker for a given
//! request context, across the policies sgl-router actually ships
//! (round-robin, random, power-of-two-choices). The cache-aware-zmq
//! policy lives in `tree_lookup.rs`; this file targets the non-tree
//! policies' steady-state hot path.
//! request context, across round-robin, random, and power-of-two choices.
use criterion::{black_box, criterion_group, criterion_main, BenchmarkId, Criterion, Throughput};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
+4 -1
View File
@@ -25,7 +25,6 @@ The dashboard graphs every family the router emits:
| `sgl_router_worker_requests_total` | Counter | Per-worker **dispatches** by `worker_url`, `model_id`, `mode`, `outcome` (recorded after dispatch; blind to pre-dispatch drops) |
| `sgl_router_request_duration_seconds` | Histogram | End-to-end request latency by `model_id` |
| `sgl_router_ttft_seconds` | Histogram | Time to first token (streaming) by `model_id` |
| `sgl_router_overlap_blocks` | Histogram | Cache-aware-zmq overlap blocks by `model_id` |
| `sgl_router_active_load` | Gauge | Per-worker prefill-token / decode-block load |
| `sgl_router_workers` | Gauge | Registered worker count by `mode` |
| `sgl_router_worker_health` | Gauge | Per-worker health (1=breaker admits, 0=open) |
@@ -35,6 +34,10 @@ The dashboard graphs every family the router emits:
| `sgl_router_decode_affinity_total` | Counter | PD decode-affinity outcomes |
| `sgl_router_sticky_total` | Counter | Sticky-session selection outcomes |
The legacy `sgl_router_overlap_blocks` metric was removed with the
`cache_aware_zmq` policy and has no direct replacement. Remove queries, alerts,
and dashboard panels that depend on this metric before upgrading.
The `sgl_router_workers` / `sgl_router_worker_*` gauges are sampled from the
live worker registry on every scrape, so a removed worker stops emitting
series immediately rather than leaving a stale value.
@@ -1605,111 +1605,6 @@
}
]
},
{
"type": "timeseries",
"title": "Overlap blocks quantiles",
"description": "Cache-aware-zmq overlap-block count at policy selection.",
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"gridPos": {
"x": 16,
"y": 41,
"w": 8,
"h": 8
},
"id": 23,
"fieldConfig": {
"defaults": {
"unit": "short",
"color": {
"mode": "palette-classic"
},
"custom": {
"axisBorderShow": false,
"axisCenteredZero": false,
"axisColorMode": "text",
"axisPlacement": "auto",
"barAlignment": 0,
"drawStyle": "line",
"fillOpacity": 0,
"gradientMode": "none",
"hideFrom": {
"legend": false,
"tooltip": false,
"viz": false
},
"insertNulls": false,
"lineInterpolation": "linear",
"lineWidth": 1,
"pointSize": 5,
"scaleDistribution": {
"type": "linear"
},
"showPoints": "never",
"spanNulls": false,
"stacking": {
"group": "A",
"mode": "none"
},
"thresholdsStyle": {
"mode": "off"
}
},
"mappings": [],
"thresholds": {
"mode": "absolute",
"steps": [
{
"color": "green",
"value": null
}
]
}
},
"overrides": []
},
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"pluginVersion": "11.6.0",
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"editorMode": "code",
"expr": "histogram_quantile(0.5, sum by (le) (rate(sgl_router_overlap_blocks_bucket{model_id=~\"$model_id\"}[$__rate_interval])))",
"range": true,
"refId": "A",
"legendFormat": "p50",
"format": "time_series"
},
{
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"editorMode": "code",
"expr": "histogram_quantile(0.99, sum by (le) (rate(sgl_router_overlap_blocks_bucket{model_id=~\"$model_id\"}[$__rate_interval])))",
"range": true,
"refId": "B",
"legendFormat": "p99",
"format": "time_series"
}
]
},
{
"type": "row",
"title": "Routing Policy",
@@ -32,4 +32,5 @@ tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt"] }
zeromq = { version = "0.6", default-features = false, features = ["tokio-runtime", "tcp-transport"] }
[build-dependencies]
protoc-bin-vendored = "3.2.0"
tonic-prost-build = "0.14.6"
@@ -133,7 +133,8 @@ sgl-router \
--model-id <model-id> \
--tokenizer-path <huggingface-repo-or-tokenizer> \
--worker-urls http://127.0.0.1:30000 \
--policy cache_aware_zmq \
--policy cache_aware \
--cache-prefix-provider indexer \
--kv-indexer-endpoint http://127.0.0.1:50051 \
--kv-indexer-query-timeout-ms 100 \
--kv-indexer-query-max-inflight 32
@@ -2,9 +2,14 @@
// SPDX-License-Identifier: Apache-2.0
fn main() -> Result<(), Box<dyn std::error::Error>> {
// Parent hashes need proto3 presence semantics to distinguish roots from
// valid hash values. Use the bundled compiler to keep codegen stable.
let mut config = tonic_prost_build::Config::new();
config.protoc_executable(protoc_bin_vendored::protoc_bin_path()?);
tonic_prost_build::configure()
.build_client(true)
.build_server(true)
.compile_protos(&["proto/kv_indexer.proto"], &["proto"])?;
.compile_with_config(config, &["proto/kv_indexer.proto"], &["proto"])?;
Ok(())
}
@@ -71,6 +71,11 @@ message ExternalKvAction {
// REPORT only. Per-hash token count (block_size), index-aligned with `hashes`,
// used to accumulate SWA trailing windows. Empty when not supplied (legacy).
repeated uint32 block_sizes = 5;
// REPORT only. Parent of hashes[0]; absent means hashes[0] is a root block.
// Every later hash is the child of the preceding hash. The Indexer and Bridge
// are deployed together, so this new protocol does not support old senders.
optional sfixed64 parent_block_hash = 6;
}
message ApplyExternalKvBatchRequest {
@@ -6,13 +6,15 @@ use std::sync::Arc;
use std::{env, io};
use sgl_kv_indexer::{
server_builder, shutdown_signal, stamp_arrival, InMemoryKvIndexerBackend, KvIndexerBackend,
KvIndexerService, DEFAULT_PREFIX_QUERY_MAX_INFLIGHT, MAX_CONCURRENT_STREAMS,
server_builder_with_max_concurrent_streams, shutdown_signal, stamp_arrival,
InMemoryKvIndexerBackend, KvIndexerBackend, KvIndexerService,
DEFAULT_PREFIX_QUERY_MAX_INFLIGHT, MAX_CONCURRENT_STREAMS,
};
use tonic::service::interceptor::InterceptedService;
use tracing::info;
const PREFIX_QUERY_MAX_INFLIGHT_ENV: &str = "KV_INDEXER_PREFIX_QUERY_MAX_INFLIGHT";
const MAX_CONCURRENT_STREAMS_ENV: &str = "KV_INDEXER_MAX_CONCURRENT_STREAMS";
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
@@ -26,6 +28,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
.unwrap_or_else(|_| "[::1]:50051".to_string())
.parse::<SocketAddr>()?;
let prefix_query_max_inflight = prefix_query_max_inflight_from_env()?;
let max_concurrent_streams = max_concurrent_streams_from_env()?;
let backend: Arc<dyn KvIndexerBackend> = Arc::new(InMemoryKvIndexerBackend::new());
// The interceptor timestamps each request before its own task is queued,
@@ -39,10 +42,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
info!(
%addr,
prefix_query_max_inflight,
max_concurrent_streams = MAX_CONCURRENT_STREAMS,
max_concurrent_streams,
"starting single-server in-memory SGLang KV Indexer"
);
server_builder()
server_builder_with_max_concurrent_streams(max_concurrent_streams)
.add_service(service)
.serve_with_shutdown(addr, shutdown_signal())
.await?;
@@ -77,6 +80,33 @@ fn parse_prefix_query_max_inflight(raw: &str) -> io::Result<usize> {
Ok(value)
}
fn max_concurrent_streams_from_env() -> io::Result<u32> {
match env::var(MAX_CONCURRENT_STREAMS_ENV) {
Ok(raw) => parse_max_concurrent_streams(&raw),
Err(env::VarError::NotPresent) => Ok(MAX_CONCURRENT_STREAMS),
Err(env::VarError::NotUnicode(_)) => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("{MAX_CONCURRENT_STREAMS_ENV} must be valid UTF-8"),
)),
}
}
fn parse_max_concurrent_streams(raw: &str) -> io::Result<u32> {
let value = raw.parse::<u32>().map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("{MAX_CONCURRENT_STREAMS_ENV} must be a positive integer, got {raw:?}"),
)
})?;
if value == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("{MAX_CONCURRENT_STREAMS_ENV} must be greater than zero"),
));
}
Ok(value)
}
#[cfg(test)]
mod tests {
use super::*;
@@ -91,4 +121,16 @@ mod tests {
assert!(parse_prefix_query_max_inflight("0").is_err());
assert!(parse_prefix_query_max_inflight("many").is_err());
}
#[test]
fn parses_positive_stream_limit() {
assert_eq!(parse_max_concurrent_streams("512").unwrap(), 512);
}
#[test]
fn rejects_invalid_stream_limit() {
assert!(parse_max_concurrent_streams("0").is_err());
assert!(parse_max_concurrent_streams("many").is_err());
assert!(parse_max_concurrent_streams("4294967296").is_err());
}
}
@@ -150,6 +150,7 @@ fn classify_rpc(status: Status) -> BridgeError {
enum Action {
Report {
tier: i32,
parent_block_hash: Option<i64>,
hashes: Vec<i64>,
masks: Vec<Option<u32>>,
block_sizes: Vec<Option<u32>>,
@@ -171,19 +172,27 @@ impl EventActions {
/// with an immediately-preceding store to the same tier and never across a
/// revoke/clear, so the final per-hash state is preserved. All hashes here
/// share the event's component mask and block size.
fn report(&mut self, tier: i32, hashes: Vec<i64>, mask: Option<u32>, block_size: Option<u32>) {
fn report(
&mut self,
tier: i32,
parent_block_hash: Option<i64>,
hashes: Vec<i64>,
mask: Option<u32>,
block_size: Option<u32>,
) {
if hashes.is_empty() {
return;
}
let n = hashes.len();
if let Some(Action::Report {
tier: last_tier,
parent_block_hash: _,
hashes: last,
masks,
block_sizes,
}) = self.actions.last_mut()
{
if *last_tier == tier {
if *last_tier == tier && parent_block_hash == last.last().copied() {
last.extend(hashes);
masks.extend(std::iter::repeat_n(mask, n));
block_sizes.extend(std::iter::repeat_n(block_size, n));
@@ -192,6 +201,7 @@ impl EventActions {
}
self.actions.push(Action::Report {
tier,
parent_block_hash,
hashes,
masks: vec![mask; n],
block_sizes: vec![block_size; n],
@@ -401,6 +411,7 @@ fn build_apply_request(
match action {
Action::Report {
tier,
parent_block_hash,
hashes,
masks,
block_sizes,
@@ -413,6 +424,7 @@ fn build_apply_request(
// the backend keeps the whole-block fast path.
component_masks: encode_component_masks(&masks),
block_sizes: encode_block_sizes(&block_sizes),
parent_block_hash,
}),
Action::Revoke { tier, hashes } => actions.push(ExternalKvAction {
r#type: ExternalKvActionType::ActionRevoke as i32,
@@ -420,6 +432,7 @@ fn build_apply_request(
hashes,
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}),
Action::ClearAll => {
for tier in &config.clear_tiers {
@@ -429,6 +442,7 @@ fn build_apply_request(
hashes: Vec::new(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
});
}
}
@@ -499,6 +513,15 @@ fn split_action(action: ExternalKvAction) -> Vec<ExternalKvAction> {
hashes: action.hashes[start..end].to_vec(),
component_masks: slice_or_empty(&action.component_masks, start, end),
block_sizes: slice_or_empty(&action.block_sizes, start, end),
parent_block_hash: if action.r#type == ExternalKvActionType::ActionReport as i32 {
if start == 0 {
action.parent_block_hash
} else {
Some(action.hashes[start - 1])
}
} else {
None
},
}
})
.collect()
@@ -599,7 +622,13 @@ fn decode_event(event: &Value, actions: &mut EventActions) -> Result<(), BridgeE
Some(_) => Some(decode_block_size(&event[4])?),
None => None,
};
actions.report(tier, decode_hashes(&event[1])?, mask, block_size);
actions.report(
tier,
decode_optional_hash(&event[2], "BlockStored.parent_block_hash")?,
decode_hashes(&event[1])?,
mask,
block_size,
);
}
"BlockRemoved" => {
if event.len() < 3 {
@@ -623,26 +652,29 @@ fn decode_event(event: &Value, actions: &mut EventActions) -> Result<(), BridgeE
fn decode_hashes(value: &Value) -> Result<Vec<i64>, BridgeError> {
expect_array(value, "block_hashes")?
.iter()
.map(|value| {
if let Some(value) = value.as_i64() {
return Ok(value);
}
// SGLang folds the unsigned top 64 bits of the SHA-256 into the
// signed range by subtracting 2^64 (`hash_str_to_int64`), which is
// two's complement, so a producer that serialises the unsigned half
// instead is carrying identical bits. Reinterpreting recovers the
// hash the router queries for; refusing the value would instead skip
// the whole event and lose every placement it carried.
if let Some(value) = value.as_u64() {
return Ok(value as i64);
}
Err(BridgeError::Decode(
"block hash must be an integer".to_string(),
))
})
.map(|value| decode_hash(value, "block hash"))
.collect()
}
fn decode_hash(value: &Value, field: &str) -> Result<i64, BridgeError> {
if let Some(value) = value.as_i64() {
return Ok(value);
}
// SGLang folds the unsigned top 64 bits of the SHA-256 into the signed
// range by subtracting 2^64. Reinterpreting recovers the same bits.
if let Some(value) = value.as_u64() {
return Ok(value as i64);
}
Err(BridgeError::Decode(format!("{field} must be an integer")))
}
fn decode_optional_hash(value: &Value, field: &str) -> Result<Option<i64>, BridgeError> {
if matches!(value, Value::Nil) {
return Ok(None);
}
decode_hash(value, field).map(Some)
}
/// Decodes the optional `component_types` slot of a `BlockStored` into a component
/// bitmask. `nil` maps to `None`, a legacy whole-block store; an array of labels
/// folds into a bitmask, and labels this build does not model are ignored.
@@ -836,10 +868,14 @@ mod tests {
}
fn stored(hashes: &[i64], medium: &str) -> Value {
stored_with_parent(hashes, None, medium)
}
fn stored_with_parent(hashes: &[i64], parent: Option<i64>, medium: &str) -> Value {
Value::Array(vec![
Value::String("BlockStored".into()),
ints(hashes),
Value::Nil, // parent_block_hash
parent.map_or(Value::Nil, Value::from),
ints(&[1]), // token_ids
Value::from(1_i64), // block_size
Value::Nil, // lora_id
@@ -850,10 +886,20 @@ mod tests {
/// A component-aware `BlockStored` (8-element schema): trailing
/// `component_types` slot plus a concrete `block_size` token count.
fn stored_c(hashes: &[i64], medium: &str, block_size: i64, components: Value) -> Value {
stored_c_with_parent(hashes, None, medium, block_size, components)
}
fn stored_c_with_parent(
hashes: &[i64],
parent: Option<i64>,
medium: &str,
block_size: i64,
components: Value,
) -> Value {
Value::Array(vec![
Value::String("BlockStored".into()),
ints(hashes),
Value::Nil, // parent_block_hash
parent.map_or(Value::Nil, Value::from),
ints(&[1]), // token_ids
Value::from(block_size),
Value::Nil, // lora_id
@@ -868,8 +914,13 @@ mod tests {
/// Legacy (whole-block) report action expectation.
fn rep(tier: i32, hashes: &[&str]) -> Action {
rep_with_parent(tier, None, hashes)
}
fn rep_with_parent(tier: i32, parent_block_hash: Option<i64>, hashes: &[&str]) -> Action {
Action::Report {
tier,
parent_block_hash,
hashes: hashes.iter().map(|h| h.parse().unwrap()).collect(),
masks: vec![None; hashes.len()],
block_sizes: vec![None; hashes.len()],
@@ -948,6 +999,7 @@ mod tests {
hashes: hashes.iter().map(|h| h.parse().unwrap()).collect(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}
}
@@ -958,6 +1010,7 @@ mod tests {
hashes: hashes.iter().map(|h| h.parse().unwrap()).collect(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}
}
@@ -968,6 +1021,7 @@ mod tests {
hashes: Vec::new(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}
}
@@ -986,6 +1040,19 @@ mod tests {
assert_eq!(request.worker_address, "127.0.0.1:9000");
}
#[test]
fn request_carries_parent_block_hash() {
let config = test_config(vec![hbm()]);
let request = request_of(
&config,
0,
vec![stored_with_parent(&[2, 3], Some(1), "GPU")],
);
assert_eq!(request.actions.len(), 1);
assert_eq!(request.actions[0].parent_block_hash, Some(1));
assert_eq!(request.actions[0].hashes, vec![2, 3]);
}
#[test]
fn oversized_report_is_split_with_aligned_metadata() {
let count = MAX_HASHES_PER_REQUEST + 1;
@@ -998,6 +1065,7 @@ mod tests {
hashes: (0..count).map(|index| index as i64).collect(),
component_masks: (0..count as u32).collect(),
block_sizes: (0..count as u32).map(|index| index + 1).collect(),
parent_block_hash: None,
}],
worker_address: "http://worker-1".into(),
cache_spec: None,
@@ -1007,6 +1075,7 @@ mod tests {
assert_eq!(batches.len(), 2);
assert_eq!(batches[0].actions[0].hashes.len(), MAX_HASHES_PER_REQUEST);
assert_eq!(batches[0].actions[0].parent_block_hash, None);
assert_eq!(
batches[1].actions[0].hashes,
vec![MAX_HASHES_PER_REQUEST as i64]
@@ -1015,6 +1084,10 @@ mod tests {
batches[1].actions[0].component_masks,
vec![MAX_HASHES_PER_REQUEST as u32]
);
assert_eq!(
batches[1].actions[0].parent_block_hash,
Some(MAX_HASHES_PER_REQUEST as i64 - 1)
);
assert_eq!(
batches[1].actions[0].block_sizes,
vec![MAX_HASHES_PER_REQUEST as u32 + 1]
@@ -1167,11 +1240,22 @@ mod tests {
#[test]
fn adjacent_same_tier_stores_coalesce() {
assert_eq!(
actions_of(vec![stored(&[1], "GPU"), stored(&[2], "GPU")]),
actions_of(vec![
stored(&[1], "GPU"),
stored_with_parent(&[2], Some(1), "GPU")
]),
vec![rep(hbm(), &["1", "2"])]
);
}
#[test]
fn same_tier_stores_on_different_chains_do_not_coalesce() {
assert_eq!(
actions_of(vec![stored(&[1], "GPU"), stored(&[2], "GPU")]),
vec![rep(hbm(), &["1"]), rep(hbm(), &["2"])]
);
}
#[test]
fn different_tier_stores_do_not_coalesce() {
assert_eq!(
@@ -1219,7 +1303,7 @@ mod tests {
assert_eq!(
decode_event_batch(&payload).unwrap().actions,
vec![
rep(hbm(), &["1234567890123", "-987654321"]),
rep_with_parent(hbm(), Some(42), &["1234567890123", "-987654321"]),
rev(ssd(), &["100", "200"]),
Action::ClearAll,
]
@@ -1352,6 +1436,7 @@ mod tests {
actions_of(vec![stored_c(&[1], "GPU", 64, strv(&["full", "swa"]))]),
vec![Action::Report {
tier: hbm(),
parent_block_hash: None,
hashes: vec![1],
masks: vec![Some(
crate::service::COMPONENT_FULL | crate::service::COMPONENT_SWA
@@ -1379,7 +1464,7 @@ mod tests {
0,
vec![
stored_c(&[1], "GPU", 64, strv(&["full", "swa"])),
stored_c(&[2], "GPU", 32, strv(&["full"])),
stored_c_with_parent(&[2], Some(1), "GPU", 32, strv(&["full"])),
],
);
assert_eq!(request.actions.len(), 1);
@@ -24,9 +24,10 @@ pub use client::{
};
pub use memory_backend::InMemoryKvIndexerBackend;
pub use service::{
component_bit, server_builder, BlockComponents, KvIndexerBackend, KvIndexerService,
WorkerPrefixInput, COMPONENT_FULL, COMPONENT_MAMBA, COMPONENT_SWA,
DEFAULT_PREFIX_QUERY_MAX_INFLIGHT, MAX_CONCURRENT_STREAMS, MAX_GRPC_DECODING_MESSAGE_SIZE,
component_bit, server_builder, server_builder_with_max_concurrent_streams, BlockComponents,
KvIndexerBackend, KvIndexerService, WorkerPrefixInput, COMPONENT_FULL, COMPONENT_MAMBA,
COMPONENT_SWA, DEFAULT_PREFIX_QUERY_MAX_INFLIGHT, MAX_CONCURRENT_STREAMS,
MAX_GRPC_DECODING_MESSAGE_SIZE,
};
pub use shutdown::shutdown_signal;
/// Re-exported because [`PrefixIndexError::Rejected`] carries it, so callers can
@@ -7,26 +7,42 @@
//! atomic and every query a consistent snapshot. The state is soft: not shared
//! with another server, and lost when the process exits.
use std::collections::{BTreeMap, HashMap, HashSet};
use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
use std::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard};
use tonic::Status;
use crate::pb::{
ApplyExternalKvBatchRequest, ApplyExternalKvBatchResponse, ExternalKvActionType,
ExternalKvNodeMatch, GetExternalKvHitCountsRequest, GetExternalKvHitCountsResponse,
HitCountEntry, MatchExternalKvPrefixRequest, MatchExternalKvPrefixResponse,
MatchExternalKvRequest, MatchExternalKvResponse, TierHashes, WorkerCacheSpec,
ApplyExternalKvBatchRequest, ApplyExternalKvBatchResponse, ExternalKvAction,
ExternalKvActionType, ExternalKvNodeMatch, GetExternalKvHitCountsRequest,
GetExternalKvHitCountsResponse, HitCountEntry, MatchExternalKvPrefixRequest,
MatchExternalKvPrefixResponse, MatchExternalKvRequest, MatchExternalKvResponse, TierHashes,
TierType, WorkerCacheSpec,
};
use crate::service::{assemble_prefix_response, prefix_limit, WorkerPrefixScanner};
use crate::service::{assemble_prefix_response, prefix_limit, WorkerPrefixScanner, COMPONENT_FULL};
use crate::{BlockComponents, KvIndexerBackend, WorkerPrefixInput};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
enum ParentLink {
#[default]
Unknown,
Root,
Hash(i64),
}
#[derive(Debug, Default)]
struct BlockRecord {
/// Chain relationship reported by the worker. Prefix-derived state is valid
/// only along links rooted at `Root`.
parent: ParentLink,
children: HashSet<i64>,
/// Shared block token count. A zero value means legacy/unspecified.
token_count: u32,
/// Resident component snapshot for each `(worker, tier)`.
placements: HashMap<(String, i32), u32>,
/// Workers for which the root-to-this-block prefix is complete and this
/// boundary is servable by the Legacy/FULL-only fast path.
prefix_complete_workers: HashSet<String>,
}
#[derive(Debug, Default)]
@@ -35,6 +51,9 @@ struct WorkerRecord {
spec: Option<WorkerCacheSpec>,
/// Reverse index used by CLEAR_ALL_AT_TIER.
holdings: HashMap<i32, HashSet<i64>>,
/// Number of non-legacy component placements. A spec-less worker can use
/// the derived fast path only while this is zero.
component_placement_count: usize,
}
#[derive(Debug, Default)]
@@ -59,6 +78,12 @@ struct PrefixCandidate {
scanner: WorkerPrefixScanner,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FastPathKind {
Legacy,
FullOnly { tier_mask: u32 },
}
/// Single-process, soft-state KV placement index.
#[derive(Debug, Default)]
pub struct InMemoryKvIndexerBackend {
@@ -89,6 +114,10 @@ impl InMemoryKvIndexerBackend {
let mut state = self.write_state()?;
let worker_id = req.worker_id;
validate_actions(&state, &req.actions)?;
let previous_fast_path = state.workers.get(&worker_id).and_then(fast_path_kind);
// Address and spec are snapshots carried on every batch. Empty address
// makes the worker unroutable; absent spec returns it to legacy mode.
{
@@ -97,16 +126,21 @@ impl InMemoryKvIndexerBackend {
worker.spec = req.cache_spec;
}
let mut dirty_roots = Vec::new();
let mut reported_chains = Vec::new();
let mut revoked_hashes = Vec::new();
// Only a fast-path worker identity change requires a full recompute.
// REPORT, REVOKE, and CLEAR enqueue their affected hashes directly.
let mut recompute_from_graph_roots = false;
for action in req.actions {
match ExternalKvActionType::try_from(action.r#type) {
Ok(ExternalKvActionType::ActionReport) => {
let has_masks = !action.component_masks.is_empty();
let has_sizes = !action.block_sizes.is_empty();
let hashes = action.hashes;
apply_report_chain(&mut state, action.parent_block_hash, &hashes);
// REPORT is a REPLACE snapshot. Keep the final occurrence
// when a coalesced action repeats one hash.
let mut last_by_hash: HashMap<i64, (u32, u32)> = HashMap::new();
for (index, hash) in action.hashes.into_iter().enumerate() {
for (index, hash) in hashes.iter().copied().enumerate() {
let mask = if has_masks {
action.component_masks[index]
} else {
@@ -117,18 +151,22 @@ impl InMemoryKvIndexerBackend {
} else {
0
};
last_by_hash.insert(hash, (mask, token_count));
}
for (hash, (mask, token_count)) in last_by_hash {
let block = state.blocks.entry(hash).or_default();
block
let old_mask = state
.blocks
.entry(hash)
.or_default()
.placements
.insert((worker_id.clone(), action.tier), mask);
adjust_component_placement_count(
&mut state,
&worker_id,
old_mask,
Some(mask),
);
// A legacy report carries no size, so 0 means
// "unknown" and must not erase a known count.
if token_count > 0 {
block.token_count = token_count;
state.blocks.entry(hash).or_default().token_count = token_count;
}
state
.workers
@@ -139,10 +177,15 @@ impl InMemoryKvIndexerBackend {
.or_default()
.insert(hash);
}
if !hashes.is_empty() {
reported_chains.push(hashes);
}
}
Ok(ExternalKvActionType::ActionRevoke) => {
for hash in action.hashes {
revoke_one(&mut state, &worker_id, &hash, action.tier);
dirty_roots.push(hash);
revoked_hashes.push(hash);
}
}
Ok(ExternalKvActionType::ActionClearAllAtTier) => {
@@ -154,6 +197,8 @@ impl InMemoryKvIndexerBackend {
.unwrap_or_default();
for hash in hashes {
revoke_one(&mut state, &worker_id, &hash, action.tier);
dirty_roots.push(hash);
revoked_hashes.push(hash);
}
}
Ok(ExternalKvActionType::ActionUnknown) | Err(_) => {
@@ -162,6 +207,31 @@ impl InMemoryKvIndexerBackend {
}
}
let current_fast_path = state.workers.get(&worker_id).and_then(fast_path_kind);
if previous_fast_path != current_fast_path {
for block in state.blocks.values_mut() {
block.prefix_complete_workers.remove(&worker_id);
}
recompute_from_graph_roots = true;
}
if recompute_from_graph_roots {
dirty_roots = state
.blocks
.iter()
.filter_map(|(hash, block)| (block.parent == ParentLink::Root).then_some(*hash))
.collect();
} else {
for hashes in reported_chains {
dirty_roots.extend(refresh_linear_report_chain_prefix_completeness(
&mut state, &worker_id, &hashes,
));
}
}
recompute_worker_subtrees(&mut state, &worker_id, dirty_roots);
for hash in revoked_hashes {
prune_empty_leaf(&mut state, hash);
}
Ok(ApplyExternalKvBatchResponse {})
}
@@ -302,11 +372,28 @@ impl InMemoryKvIndexerBackend {
blocks_read: 1,
});
};
let known_prefix_len = known_request_prefix_len(&state, hashes).unwrap_or(0);
let fast_worker_ids: HashSet<&str> = if known_prefix_len > 0 {
first
.prefix_complete_workers
.iter()
.filter_map(|worker_id| {
state.workers.get(worker_id).and_then(|worker| {
(fast_path_kind(worker).is_some() && !worker.address.is_empty())
.then_some(worker_id.as_str())
})
})
.collect()
} else {
HashSet::new()
};
let mut seen = HashSet::new();
let mut candidates: Vec<PrefixCandidate> = first
.placements
.keys()
.filter(|(worker, _)| seen.insert(worker.as_str()))
.filter(|(worker, _)| {
!fast_worker_ids.contains(worker.as_str()) && seen.insert(worker.as_str())
})
.map(|(worker, _)| {
let metadata = state.workers.get(worker);
PrefixCandidate {
@@ -332,41 +419,61 @@ impl InMemoryKvIndexerBackend {
tier_masks: Vec::new(),
})
.collect();
for hash in hashes {
present.fill(false);
for block in &mut block_views {
block.token_count = 0;
block.tier_masks.clear();
let mut entries = Vec::with_capacity(fast_worker_ids.len() + candidates.len());
let mut unresolved = fast_worker_ids;
for (index, hash) in hashes[..known_prefix_len].iter().enumerate().rev() {
if unresolved.is_empty() {
break;
}
if let Some(block) = state.blocks.get(hash) {
for ((worker, tier), mask) in &block.placements {
let Some(&index) = candidate_by_id.get(worker) else {
for worker_id in &block.prefix_complete_workers {
if !unresolved.remove(worker_id.as_str()) {
continue;
};
present[index] = true;
block_views[index].token_count = block.token_count;
block_views[index].tier_masks.push((*tier, *mask));
}
if let Some(worker) = state.workers.get(worker_id) {
entries.push((
worker_id.clone(),
worker.address.clone(),
(index + 1) as u32,
));
}
}
}
for (index, candidate) in candidates.iter_mut().enumerate() {
candidate
.scanner
.push(present[index].then_some(&block_views[index]));
}
}
let entries = candidates
.into_iter()
.filter_map(|candidate| {
let prefix = candidate.scanner.prefix();
(!candidate.address.is_empty() && prefix > 0).then_some((
candidate.worker_id,
candidate.address,
prefix,
))
})
.collect();
if !candidates.is_empty() {
for hash in hashes {
present.fill(false);
for block in &mut block_views {
block.token_count = 0;
block.tier_masks.clear();
}
if let Some(block) = state.blocks.get(hash) {
for ((worker, tier), mask) in &block.placements {
let Some(&index) = candidate_by_id.get(worker) else {
continue;
};
present[index] = true;
block_views[index].token_count = block.token_count;
block_views[index].tier_masks.push((*tier, *mask));
}
}
for (index, candidate) in candidates.iter_mut().enumerate() {
candidate
.scanner
.push(present[index].then_some(&block_views[index]));
}
}
}
entries.extend(candidates.into_iter().filter_map(|candidate| {
let prefix = candidate.scanner.prefix();
(!candidate.address.is_empty() && prefix > 0).then_some((
candidate.worker_id,
candidate.address,
prefix,
))
}));
Ok(assemble_prefix_response(entries, limit as u32))
}
@@ -392,12 +499,319 @@ impl InMemoryKvIndexerBackend {
}
}
fn revoke_one(state: &mut State, worker_id: &str, hash: &i64, tier: i32) {
let mut remove_block = false;
if let Some(block) = state.blocks.get_mut(hash) {
block.placements.remove(&(worker_id.to_string(), tier));
remove_block = block.placements.is_empty();
fn fast_path_kind(worker: &WorkerRecord) -> Option<FastPathKind> {
match worker.spec.as_ref() {
None if worker.component_placement_count == 0 => Some(FastPathKind::Legacy),
Some(spec) if spec.version <= 1 && spec.components == COMPONENT_FULL => {
Some(FastPathKind::FullOnly {
tier_mask: spec.full_tier_mask,
})
}
_ => None,
}
}
fn tier_in_mask(mask: u32, tier: i32) -> bool {
tier >= 0 && mask & (1u32 << tier) != 0
}
fn block_servable(state: &State, hash: i64, worker_id: &str, kind: FastPathKind) -> bool {
match kind {
// A globally component-free legacy worker only needs membership, which
// its reverse holdings index answers without scanning every other
// worker placed on this popular block.
FastPathKind::Legacy => state.workers.get(worker_id).is_some_and(|worker| {
worker
.holdings
.values()
.any(|hashes| hashes.contains(&hash))
}),
FastPathKind::FullOnly { tier_mask } => {
let indexer_tiers =
(1 << (TierType::TierHbm as u32)) | (1 << (TierType::TierDram as u32));
state.blocks.get(&hash).is_some_and(|block| {
block.placements.iter().any(|((worker, tier), mask)| {
worker == worker_id
&& mask & COMPONENT_FULL != 0
&& tier_in_mask(indexer_tiers & tier_mask, *tier)
})
})
}
}
}
#[cfg(test)]
fn link_report_chain(
state: &mut State,
parent_block_hash: Option<i64>,
hashes: &[i64],
) -> Result<(), Status> {
let mut planned_parents = HashMap::with_capacity(hashes.len());
validate_report_chain(state, &mut planned_parents, parent_block_hash, hashes)?;
validate_parent_graph_acyclic(state, &planned_parents)?;
apply_report_chain(state, parent_block_hash, hashes);
Ok(())
}
fn apply_report_chain(state: &mut State, parent_block_hash: Option<i64>, hashes: &[i64]) {
let mut parent = parent_block_hash.map_or(ParentLink::Root, ParentLink::Hash);
for hash in hashes {
if let ParentLink::Hash(parent_hash) = parent {
state
.blocks
.entry(parent_hash)
.or_default()
.children
.insert(*hash);
}
state.blocks.entry(*hash).or_default().parent = parent;
parent = ParentLink::Hash(*hash);
}
}
fn validate_actions(state: &State, actions: &[ExternalKvAction]) -> Result<(), Status> {
let mut planned_parents = HashMap::new();
for action in actions {
match ExternalKvActionType::try_from(action.r#type) {
Ok(ExternalKvActionType::ActionReport) => validate_report_chain(
state,
&mut planned_parents,
action.parent_block_hash,
&action.hashes,
)?,
Ok(ExternalKvActionType::ActionRevoke)
| Ok(ExternalKvActionType::ActionClearAllAtTier) => {}
Ok(ExternalKvActionType::ActionUnknown) | Err(_) => {
return Err(Status::invalid_argument("unsupported action type"));
}
}
}
validate_parent_graph_acyclic(state, &planned_parents)
}
fn validate_report_chain(
state: &State,
planned_parents: &mut HashMap<i64, ParentLink>,
parent_block_hash: Option<i64>,
hashes: &[i64],
) -> Result<(), Status> {
let mut parent = parent_block_hash.map_or(ParentLink::Root, ParentLink::Hash);
for hash in hashes {
if parent == ParentLink::Hash(*hash) {
return Err(Status::invalid_argument(
"block hash cannot be its own parent",
));
}
let existing = planned_parents
.get(hash)
.copied()
.or_else(|| state.blocks.get(hash).map(|block| block.parent))
.unwrap_or_default();
if existing != ParentLink::Unknown && existing != parent {
return Err(Status::invalid_argument(format!(
"block hash {hash} was reported with conflicting parents"
)));
}
planned_parents.insert(*hash, parent);
parent = ParentLink::Hash(*hash);
}
Ok(())
}
fn validate_parent_graph_acyclic(
state: &State,
planned_parents: &HashMap<i64, ParentLink>,
) -> Result<(), Status> {
let mut complete = HashSet::new();
for start in planned_parents.keys().copied() {
if complete.contains(&start) {
continue;
}
let mut path = Vec::new();
let mut on_path = HashSet::new();
let mut current = start;
loop {
if complete.contains(&current) {
break;
}
if !on_path.insert(current) {
return Err(Status::invalid_argument(
"report would create a parent cycle",
));
}
path.push(current);
let parent = planned_parents
.get(&current)
.copied()
.or_else(|| state.blocks.get(&current).map(|block| block.parent))
.unwrap_or_default();
match parent {
ParentLink::Hash(parent) => current = parent,
ParentLink::Unknown | ParentLink::Root => break,
}
}
complete.extend(path);
}
Ok(())
}
fn adjust_component_placement_count(
state: &mut State,
worker_id: &str,
old_mask: Option<u32>,
new_mask: Option<u32>,
) {
let worker = state.workers.entry(worker_id.to_string()).or_default();
if old_mask.is_some_and(|mask| mask != 0) {
worker.component_placement_count = worker.component_placement_count.saturating_sub(1);
}
if new_mask.is_some_and(|mask| mask != 0) {
worker.component_placement_count = worker.component_placement_count.saturating_add(1);
}
}
fn recompute_worker_subtrees(
state: &mut State,
worker_id: &str,
roots: impl IntoIterator<Item = i64>,
) {
let kind = state.workers.get(worker_id).and_then(fast_path_kind);
let mut queue: VecDeque<i64> = roots.into_iter().collect();
let mut visited = HashSet::new();
while let Some(hash) = queue.pop_front() {
if !visited.insert(hash) {
continue;
}
let Some(block) = state.blocks.get(&hash) else {
continue;
};
let parent_complete = match block.parent {
ParentLink::Unknown => false,
ParentLink::Root => true,
ParentLink::Hash(parent) => state
.blocks
.get(&parent)
.is_some_and(|parent| parent.prefix_complete_workers.contains(worker_id)),
};
let complete = kind
.is_some_and(|kind| parent_complete && block_servable(state, hash, worker_id, kind));
let children: Vec<i64> = block.children.iter().copied().collect();
let block = state.blocks.get_mut(&hash).expect("block exists");
if complete {
block.prefix_complete_workers.insert(worker_id.to_string());
} else {
block.prefix_complete_workers.remove(worker_id);
}
queue.extend(children);
}
}
/// Returns direct children held by this worker but outside the current REPORT chain.
fn external_children_held_by_worker(
state: &State,
worker_id: &str,
reported_hashes: &HashSet<i64>,
parent: i64,
) -> Vec<i64> {
state
.blocks
.get(&parent)
.into_iter()
.flat_map(|block| block.children.iter().copied())
.filter(|child| {
!reported_hashes.contains(child)
&& state.blocks.get(child).is_some_and(|child| {
child
.placements
.keys()
.any(|(worker, _)| worker == worker_id)
})
})
.collect()
}
/// Refreshes derived prefix state along a closed linear REPORT chain.
///
/// The caller has verified that no node in the chain has an external child.
fn refresh_linear_report_chain_prefix_completeness(
state: &mut State,
worker_id: &str,
hashes: &[i64],
) -> Vec<i64> {
let kind = state.workers.get(worker_id).and_then(fast_path_kind);
let reported_hashes: HashSet<i64> = hashes.iter().copied().collect();
let mut external_dirty_roots = Vec::new();
let mut parent_complete = hashes
.first()
.and_then(|hash| state.blocks.get(hash))
.is_some_and(|block| match block.parent {
ParentLink::Root => true,
ParentLink::Hash(parent) => state
.blocks
.get(&parent)
.is_some_and(|parent| parent.prefix_complete_workers.contains(worker_id)),
ParentLink::Unknown => false,
});
for hash in hashes {
let was_complete = state
.blocks
.get(hash)
.is_some_and(|block| block.prefix_complete_workers.contains(worker_id));
let complete = kind
.is_some_and(|kind| parent_complete && block_servable(state, *hash, worker_id, kind));
if was_complete != complete {
external_dirty_roots.extend(external_children_held_by_worker(
state,
worker_id,
&reported_hashes,
*hash,
));
}
let Some(block) = state.blocks.get_mut(hash) else {
continue;
};
if complete {
block.prefix_complete_workers.insert(worker_id.to_string());
} else {
block.prefix_complete_workers.remove(worker_id);
}
parent_complete = complete;
}
external_dirty_roots
}
/// Returns the length of the longest leading request chain already known to the
/// Indexer. A missing block starts the normal uncached suffix; a present block
/// with the wrong parent is a chain conflict and disables the derived fast path.
fn known_request_prefix_len(state: &State, hashes: &[i64]) -> Option<usize> {
let mut known = 0;
for (index, hash) in hashes.iter().enumerate() {
let expected = if index == 0 {
ParentLink::Root
} else {
ParentLink::Hash(hashes[index - 1])
};
let Some(block) = state.blocks.get(hash) else {
break;
};
if block.parent != expected {
return None;
}
known += 1;
}
Some(known)
}
fn revoke_one(state: &mut State, worker_id: &str, hash: &i64, tier: i32) {
let mut removed_mask = None;
if let Some(block) = state.blocks.get_mut(hash) {
removed_mask = block.placements.remove(&(worker_id.to_string(), tier));
}
adjust_component_placement_count(state, worker_id, removed_mask, None);
if let Some(worker) = state.workers.get_mut(worker_id) {
if let Some(hashes) = worker.holdings.get_mut(&tier) {
@@ -408,12 +822,38 @@ fn revoke_one(state: &mut State, worker_id: &str, hash: &i64, tier: i32) {
}
}
if remove_block {
state.blocks.remove(hash);
if state
.blocks
.get(hash)
.is_some_and(|block| block.placements.is_empty())
{
state.hit_counts.remove(hash);
}
}
fn prune_empty_leaf(state: &mut State, mut hash: i64) {
loop {
let Some(block) = state.blocks.get(&hash) else {
return;
};
if !block.placements.is_empty()
|| !block.children.is_empty()
|| !block.prefix_complete_workers.is_empty()
{
return;
}
let parent = block.parent;
state.blocks.remove(&hash);
let ParentLink::Hash(parent_hash) = parent else {
return;
};
if let Some(parent) = state.blocks.get_mut(&parent_hash) {
parent.children.remove(&hash);
}
hash = parent_hash;
}
}
fn dedup_preserve_order(hashes: &[i64]) -> Vec<i64> {
let mut seen = HashSet::new();
hashes
@@ -497,4 +937,74 @@ mod tests {
drop(read_guard);
query.join().unwrap();
}
#[test]
fn known_request_prefix_stops_at_uncached_suffix_and_rejects_conflicts() {
let mut state = State::default();
link_report_chain(&mut state, None, &[1, 2, 3]).unwrap();
link_report_chain(&mut state, None, &[9]).unwrap();
assert_eq!(known_request_prefix_len(&state, &[1, 2, 3, 4, 5]), Some(3));
assert_eq!(known_request_prefix_len(&state, &[1, 9]), None);
}
#[test]
fn conflicting_report_chain_does_not_mutate_the_graph() {
let mut state = State::default();
let error = link_report_chain(&mut state, None, &[1, 2, 1]).unwrap_err();
assert_eq!(error.code(), tonic::Code::InvalidArgument);
assert!(state.blocks.is_empty());
}
#[test]
fn cyclic_report_chain_does_not_mutate_the_graph() {
let mut state = State::default();
let error = link_report_chain(&mut state, Some(2), &[1, 2]).unwrap_err();
assert_eq!(error.code(), tonic::Code::InvalidArgument);
assert!(state.blocks.is_empty());
}
#[test]
fn cyclic_report_chain_through_existing_graph_is_rejected() {
let mut state = State::default();
link_report_chain(&mut state, Some(2), &[1]).unwrap();
let error = link_report_chain(&mut state, Some(1), &[2]).unwrap_err();
assert_eq!(error.code(), tonic::Code::InvalidArgument);
assert_eq!(state.blocks[&1].parent, ParentLink::Hash(2));
assert_eq!(state.blocks[&2].parent, ParentLink::Unknown);
}
#[test]
fn external_children_only_include_the_reporting_workers_branch() {
let mut state = State::default();
link_report_chain(&mut state, None, &[1, 2, 3]).unwrap();
link_report_chain(&mut state, Some(1), &[4]).unwrap();
state
.blocks
.get_mut(&4)
.unwrap()
.placements
.insert(("worker-b".into(), TierType::TierHbm as i32), 0);
let reported_hashes: HashSet<i64> = [1, 2, 3].into_iter().collect();
assert!(
external_children_held_by_worker(&state, "worker-a", &reported_hashes, 1,).is_empty()
);
state
.blocks
.get_mut(&4)
.unwrap()
.placements
.insert(("worker-a".into(), TierType::TierHbm as i32), 0);
assert_eq!(
external_children_held_by_worker(&state, "worker-a", &reported_hashes, 1,),
vec![4]
);
}
}
@@ -182,7 +182,16 @@ where
/// [`KvIndexerService::into_server`]: that sets the per-message ceiling, this
/// bounds how many messages can be in flight against it at once.
pub fn server_builder() -> Server {
Server::builder().max_concurrent_streams(MAX_CONCURRENT_STREAMS)
server_builder_with_max_concurrent_streams(MAX_CONCURRENT_STREAMS)
}
/// A transport builder with an explicit stream bound for high-fanout fleets.
///
/// The default entry point keeps the stable 64-stream behavior, while the
/// standalone Indexer binary can raise the bound when it has one bridge per
/// worker.
pub fn server_builder_with_max_concurrent_streams(max_concurrent_streams: u32) -> Server {
Server::builder().max_concurrent_streams(max_concurrent_streams)
}
#[tonic::async_trait]
@@ -775,6 +784,7 @@ mod tests {
hashes: hashes.iter().map(|h| h.parse().unwrap()).collect(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}
}
@@ -14,12 +14,22 @@ pub fn dram() -> i32 {
}
pub fn action(kind: ExternalKvActionType, tier: i32, hashes: &[i64]) -> ExternalKvAction {
action_with_parent(kind, tier, None, hashes)
}
pub fn action_with_parent(
kind: ExternalKvActionType,
tier: i32,
parent_block_hash: Option<i64>,
hashes: &[i64],
) -> ExternalKvAction {
ExternalKvAction {
r#type: kind as i32,
tier,
hashes: hashes.to_vec(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash,
}
}
@@ -31,6 +41,17 @@ pub fn component_report(
hashes: &[i64],
masks: &[u32],
block_sizes: &[u32],
) -> ExternalKvAction {
component_report_with_parent(tier, None, hashes, masks, block_sizes)
}
#[allow(dead_code)]
pub fn component_report_with_parent(
tier: i32,
parent_block_hash: Option<i64>,
hashes: &[i64],
masks: &[u32],
block_sizes: &[u32],
) -> ExternalKvAction {
ExternalKvAction {
r#type: ExternalKvActionType::ActionReport as i32,
@@ -38,6 +59,7 @@ pub fn component_report(
hashes: hashes.to_vec(),
component_masks: masks.to_vec(),
block_sizes: block_sizes.to_vec(),
parent_block_hash,
}
}
@@ -34,7 +34,7 @@ use sgl_kv_indexer::{
PrefixIndex, PrefixIndexConfig, MAX_GRPC_DECODING_MESSAGE_SIZE,
};
use test_id::nanos;
use test_kv::{action, apply_request, hbm};
use test_kv::{action, action_with_parent, apply_request, hbm};
use test_net::free_addr;
async fn start_backend(
@@ -211,6 +211,27 @@ fn apply_report(
)
}
fn apply_report_with_parent(
worker: &str,
addr: &str,
seq: u64,
tier: i32,
parent_block_hash: Option<i64>,
hashes: &[i64],
) -> ApplyExternalKvBatchRequest {
apply_request(
worker,
addr,
seq,
vec![action_with_parent(
ExternalKvActionType::ActionReport,
tier,
parent_block_hash,
hashes,
)],
)
}
#[tokio::test]
async fn multiple_workers_share_one_indexer_server() {
let mut indexer = start().await;
@@ -225,7 +246,7 @@ async fn multiple_workers_share_one_indexer_server() {
"10.0.0.1:9000",
1,
hbm(),
&[hash_0, shared_hash],
&[shared_hash, hash_0],
))
.await
.expect("apply worker-0");
@@ -235,7 +256,7 @@ async fn multiple_workers_share_one_indexer_server() {
"10.0.0.2:9000",
1,
hbm(),
&[hash_1, shared_hash],
&[shared_hash, hash_1],
))
.await
.expect("apply worker-1");
@@ -309,6 +330,7 @@ async fn validation_errors_map_to_invalid_argument_over_grpc() {
hashes: vec![1],
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}],
};
let err = c
@@ -318,6 +340,76 @@ async fn validation_errors_map_to_invalid_argument_over_grpc() {
assert_eq!(err.code(), Code::InvalidArgument);
}
#[tokio::test]
async fn rejected_batch_is_atomic_over_grpc() {
let mut c = start().await;
c.apply_external_kv_batch(apply_report("w", "old-address", 1, hbm(), &[1, 2]))
.await
.expect("seed chain");
let err = c
.apply_external_kv_batch(apply_request(
"w",
"new-address",
2,
vec![
action(ExternalKvActionType::ActionReport, hbm(), &[3]),
action_with_parent(ExternalKvActionType::ActionReport, hbm(), Some(9), &[2]),
],
))
.await
.expect_err("conflicting parent must reject the whole batch");
assert_eq!(err.code(), Code::InvalidArgument);
let old = c
.match_external_kv(MatchExternalKvRequest {
hashes: vec![1],
count_as_hit: false,
})
.await
.expect("query original state")
.into_inner();
assert_eq!(old.matches.len(), 1);
assert_eq!(old.matches[0].address, "old-address");
let leaked = c
.match_external_kv(MatchExternalKvRequest {
hashes: vec![3],
count_as_hit: false,
})
.await
.expect("query rejected action")
.into_inner();
assert!(leaked.matches.is_empty());
}
#[tokio::test]
async fn cyclic_report_is_rejected_over_grpc() {
let mut c = start().await;
let err = c
.apply_external_kv_batch(apply_report_with_parent(
"w",
"address",
1,
hbm(),
Some(2),
&[1, 2],
))
.await
.expect_err("cyclic report must be rejected");
assert_eq!(err.code(), Code::InvalidArgument);
let response = c
.match_external_kv(MatchExternalKvRequest {
hashes: vec![1, 2],
count_as_hit: false,
})
.await
.expect("query rejected report")
.into_inner();
assert!(response.matches.is_empty());
}
#[tokio::test]
async fn match_prefix_over_grpc() {
let mut c = start().await;
@@ -358,12 +450,14 @@ async fn prefix_query_scans_more_than_one_apply_chunk_over_grpc() {
let mut indexer = start().await;
let hashes: Vec<i64> = (0..=APPLY_CHUNK_SIZE as i64).collect();
for (seq, chunk) in hashes.chunks(APPLY_CHUNK_SIZE).enumerate() {
let parent_block_hash = (seq > 0).then_some(chunk[0] - 1);
indexer
.apply_external_kv_batch(apply_report(
.apply_external_kv_batch(apply_report_with_parent(
"large-prefix-worker",
"10.0.0.1:9000",
seq as u64,
hbm(),
parent_block_hash,
chunk,
))
.await
@@ -487,7 +581,7 @@ async fn start_recording_deadlines(
/// the only thing letting the indexer shed a query whose caller gave up.
#[tokio::test]
async fn router_client_publishes_its_deadline_on_the_wire() {
let (index, seen) = start_recording_deadlines(Duration::from_millis(250)).await;
let (index, seen) = start_recording_deadlines(Duration::from_secs(2)).await;
index
.match_prefix(vec![1, 2, 3])
@@ -17,7 +17,10 @@ use sgl_kv_indexer::pb::{
use sgl_kv_indexer::{
InMemoryKvIndexerBackend, KvIndexerBackend, WorkerPrefixInput, COMPONENT_FULL, COMPONENT_SWA,
};
use test_kv::{action, apply_request as apply_req, component_report, dram, hbm};
use test_kv::{
action, action_with_parent, apply_request as apply_req, component_report,
component_report_with_parent, dram, hbm,
};
use tonic::Status;
fn backend() -> InMemoryKvIndexerBackend {
@@ -81,6 +84,38 @@ itest!(report_then_match_returns_worker_and_address, b, {
assert!(tiers_for(&resp, "w1", 3).is_empty());
});
itest!(rejected_batch_does_not_publish_earlier_actions, b, {
b.apply_external_kv_batch(apply_req(
"w1",
"old-address",
1,
vec![action(ExternalKvActionType::ActionReport, hbm(), &[1, 2])],
))
.await
.unwrap();
let error = b
.apply_external_kv_batch(apply_req(
"w1",
"new-address",
2,
vec![
action(ExternalKvActionType::ActionReport, hbm(), &[3]),
action_with_parent(ExternalKvActionType::ActionReport, hbm(), Some(9), &[2]),
],
))
.await
.unwrap_err();
assert_eq!(error.code(), tonic::Code::InvalidArgument);
let old = b.match_external_kv(match_req(&[1], false)).await.unwrap();
assert_eq!(old.matches.len(), 1);
assert_eq!(old.matches[0].address, "old-address");
let leaked = b.match_external_kv(match_req(&[3], false)).await.unwrap();
assert!(leaked.matches.is_empty());
});
itest!(large_request_preserves_complete_ordered_results, b, {
// Exercise a large write and read while preserving complete ordered results.
let expected_hashes: Vec<i64> = (0..300).collect();
@@ -587,10 +622,19 @@ async fn prefix_fast_path_matches_default_impl() {
fast.apply_external_kv_batch(report("w-short", "10.0.0.2:1", 1, &[1, 2]))
.await
.unwrap();
// w-hole holds 1, 3, 4 but not 2: strict prefix must be 1.
fast.apply_external_kv_batch(report("w-hole", "10.0.0.3:1", 1, &[1, 3, 4]))
// w-hole first learns the same chain, then loses block 2 while descendants
// remain placed: strict prefix must be 1.
fast.apply_external_kv_batch(report("w-hole", "10.0.0.3:1", 1, &[1, 2, 3, 4]))
.await
.unwrap();
fast.apply_external_kv_batch(apply_req(
"w-hole",
"10.0.0.3:1",
2,
vec![action(ExternalKvActionType::ActionRevoke, hbm(), &[2])],
))
.await
.unwrap();
// w-noaddr is unroutable and must be excluded by both paths.
fast.apply_external_kv_batch(report("w-noaddr", "", 1, &[1, 2]))
.await
@@ -626,6 +670,64 @@ async fn prefix_fast_path_matches_default_impl() {
assert!(fast_resp.blocks_read >= 1);
}
#[tokio::test]
async fn prefix_fast_path_returns_worker_depths_with_uncached_suffix() {
let (fast, reference) = shared_state_pair();
fast.apply_external_kv_batch(report("w-long", "10.0.0.1:1", 1, &[1, 2, 3]))
.await
.unwrap();
fast.apply_external_kv_batch(report("w-short", "10.0.0.2:1", 1, &[1, 2]))
.await
.unwrap();
// Blocks 4 and 5 are the newly appended turn and are not cached anywhere.
// They must cap the maximum prefix without disabling the known-prefix path.
let query = [1, 2, 3, 4, 5];
let fast_response = fast
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
let reference_response = reference
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
assert_eq!(
prefix_pairs(&fast_response),
prefix_pairs(&reference_response)
);
assert_eq!(
prefix_pairs(&fast_response),
vec![("w-long".to_string(), 3), ("w-short".to_string(), 2)]
);
}
#[tokio::test]
async fn prefix_fast_path_falls_back_on_existing_parent_conflict() {
let (fast, reference) = shared_state_pair();
fast.apply_external_kv_batch(report("w1", "10.0.0.1:1", 1, &[1, 2]))
.await
.unwrap();
// Hash 9 is an independent root, not a child of hash 1.
fast.apply_external_kv_batch(report("w1", "10.0.0.1:1", 2, &[9]))
.await
.unwrap();
let query = [1, 9];
let fast_response = fast
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
let reference_response = reference
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
assert_eq!(
prefix_pairs(&fast_response),
prefix_pairs(&reference_response)
);
}
#[tokio::test]
async fn prefix_first_block_miss_reads_one_block() {
let b = backend();
@@ -662,6 +764,156 @@ async fn prefix_max_blocks_caps_the_scan() {
assert_eq!(resp.matches[0].matched_prefix_blocks, 2);
}
#[tokio::test]
async fn prefix_complete_revoke_and_restore_propagates_to_descendants() {
let (fast, reference) = shared_state_pair();
fast.apply_external_kv_batch(report("w1", "10.0.0.1:1", 1, &[1, 2, 3, 4]))
.await
.unwrap();
fast.apply_external_kv_batch(apply_req(
"w1",
"10.0.0.1:1",
2,
vec![action(ExternalKvActionType::ActionRevoke, hbm(), &[2])],
))
.await
.unwrap();
let after_revoke = fast
.match_external_kv_prefix(prefix_req(&[1, 2, 3, 4]))
.await
.unwrap();
assert_eq!(prefix_pairs(&after_revoke), vec![("w1".to_string(), 1)]);
fast.apply_external_kv_batch(apply_req(
"w1",
"10.0.0.1:1",
3,
vec![action_with_parent(
ExternalKvActionType::ActionReport,
hbm(),
Some(1),
&[2],
)],
))
.await
.unwrap();
let restored = fast
.match_external_kv_prefix(prefix_req(&[1, 2, 3, 4]))
.await
.unwrap();
let expected = reference
.match_external_kv_prefix(prefix_req(&[1, 2, 3, 4]))
.await
.unwrap();
assert_eq!(prefix_pairs(&restored), vec![("w1".to_string(), 4)]);
assert_eq!(prefix_pairs(&restored), prefix_pairs(&expected));
}
#[tokio::test]
async fn prefix_complete_fast_path_preserves_cache_hit_rate() {
let (fast, reference) = shared_state_pair();
let query: Vec<i64> = (1..=64).collect();
let worker_count = 32usize;
let mut expected_prefix_sum = 0u64;
for worker in 0..worker_count {
let depth = 1 + (worker * 7 % query.len());
expected_prefix_sum += depth as u64;
fast.apply_external_kv_batch(report(
&format!("w-{worker:02}"),
&format!("http://worker-{worker:02}"),
1,
&query[..depth],
))
.await
.unwrap();
}
let fast_response = fast
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
let reference_response = reference
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
assert_eq!(
prefix_pairs(&fast_response),
prefix_pairs(&reference_response)
);
assert_eq!(fast_response.matches.len(), worker_count);
let fast_prefix_sum: u64 = fast_response
.matches
.iter()
.map(|item| item.matched_prefix_blocks as u64)
.sum();
let reference_prefix_sum: u64 = reference_response
.matches
.iter()
.map(|item| item.matched_prefix_blocks as u64)
.sum();
assert_eq!(fast_prefix_sum, expected_prefix_sum);
assert_eq!(fast_prefix_sum, reference_prefix_sum);
let hit_rate = fast_prefix_sum as f64 / (worker_count * query.len()) as f64;
let reference_hit_rate = reference_prefix_sum as f64 / (worker_count * query.len()) as f64;
assert!((hit_rate - reference_hit_rate).abs() < f64::EPSILON);
}
#[tokio::test]
async fn full_only_prefix_complete_fast_path_matches_reference() {
let (fast, reference) = shared_state_pair();
let full_spec = WorkerCacheSpec {
version: 1,
components: COMPONENT_FULL,
swa_window_tokens: 0,
full_tier_mask: (1 << hbm()) | (1 << dram()),
swa_tier_mask: 0,
mamba_tier_mask: 0,
};
fast.apply_external_kv_batch(apply_with_spec(
"w-full",
"10.0.0.1:1",
1,
full_spec,
vec![component_report(
hbm(),
&[11, 12, 13, 14],
&[COMPONENT_FULL; 4],
&[16; 4],
)],
))
.await
.unwrap();
fast.apply_external_kv_batch(apply_with_spec(
"w-full",
"10.0.0.1:1",
2,
full_spec,
vec![action(ExternalKvActionType::ActionRevoke, hbm(), &[13])],
))
.await
.unwrap();
let query = [11, 12, 13, 14];
let fast_response = fast
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
let reference_response = reference
.match_external_kv_prefix(prefix_req(&query))
.await
.unwrap();
assert_eq!(
prefix_pairs(&fast_response),
prefix_pairs(&reference_response)
);
assert_eq!(
prefix_pairs(&fast_response),
vec![("w-full".to_string(), 2)]
);
}
// --- component-aware placement & prefix -------------------------------------
/// A hybrid-SWA spec: full servable from HBM+DRAM, swa a 100-token trailing
@@ -766,7 +1018,13 @@ async fn partial_eviction_replace_shrinks_component_set() {
"10.0.0.1:1",
2,
swa_spec(),
vec![component_report(hbm(), &[2], &[COMPONENT_FULL], &[80])],
vec![component_report_with_parent(
hbm(),
Some(1),
&[2],
&[COMPONENT_FULL],
&[80],
)],
))
.await
.unwrap();
@@ -827,12 +1085,10 @@ async fn duplicate_hash_in_one_report_keeps_last_snapshot() {
"10.0.0.1:1",
1,
swa_spec(),
vec![component_report(
hbm(),
&[1, 1],
&[COMPONENT_FULL | COMPONENT_SWA, COMPONENT_FULL],
&[80, 80],
)],
vec![
component_report(hbm(), &[1], &[COMPONENT_FULL | COMPONENT_SWA], &[80]),
component_report(hbm(), &[1], &[COMPONENT_FULL], &[80]),
],
))
.await
.unwrap();
+221 -137
View File
@@ -12,10 +12,10 @@ use std::num::NonZeroU32;
use crate::config::{
default_cb_cool_down, default_proxy_request_timeout_secs, default_stale_request_timeout_secs,
resolve_mode, ActiveLoadConfig, AffinityConfig, AffinityMode, CacheAwareConfig,
CircuitBreakerConfig, Config, DiscoveryBackend, EligibilityConfig, FilterKind, FusedTerm,
K8sDiscoveryConfig, KvIndexerEndpointConfig, LogFormat, ModelConfig, ObservabilityConfig,
PolicyKind, ProxyConfig, ServerConfig, SessionAffinityMode, StaticUrlsDiscoveryConfig,
StickyConfig, StickyFallbackKind, DEFAULT_FUSE,
CachePrefixProvider, CircuitBreakerConfig, Config, DecodePolicyKind, DiscoveryBackend,
EligibilityConfig, FilterKind, FusedTerm, K8sDiscoveryConfig, KvIndexerEndpointConfig,
LogFormat, ModelConfig, ObservabilityConfig, PolicyKind, ProxyConfig, ServerConfig,
SessionAffinityMode, StaticUrlsDiscoveryConfig, StickyConfig, StickyFallbackKind, DEFAULT_FUSE,
};
const DEFAULT_KV_INDEXER_QUERY_TIMEOUT_MS: u64 = 100;
@@ -53,6 +53,12 @@ pub struct Cli {
/// Routing policy.
#[arg(long, value_enum, default_value = "round_robin")]
pub policy: PolicyKind,
/// Policy used to select decode workers for PD requests.
#[arg(long, value_enum, default_value = "power_of_two")]
pub decode_policy: DecodePolicyKind,
/// Static P/D bucket configuration. Omit to use the global candidate domain.
#[arg(long)]
pub bucket_config: Option<String>,
// ---- circuit breaker (opt-in via --cb-threshold) ----
/// Consecutive upstream failures before the circuit breaker opens.
@@ -64,16 +70,6 @@ pub struct Cli {
#[arg(long)]
pub cb_cool_down_secs: Option<u64>,
// ---- legacy cache-aware-zmq tuning ----
/// Min `matched_blocks / total_blocks` for a cache match to win.
#[arg(long)]
pub cache_threshold: Option<f32>,
/// Absolute load spread above which the cache check is skipped.
#[arg(long)]
pub balance_abs_threshold: Option<usize>,
/// Multiplicative load spread gating the absolute balance check.
#[arg(long)]
pub balance_rel_threshold: Option<f32>,
/// External KV indexer gRPC endpoint used as the authoritative cache signal.
/// Needs an explicit scheme, e.g. `http://10.0.0.1:50051`.
#[arg(long)]
@@ -86,6 +82,9 @@ pub struct Cli {
/// `--kv-indexer-endpoint`; defaults to 32.
#[arg(long)]
pub kv_indexer_query_max_inflight: Option<usize>,
/// Prefix-match source for native Cache-Aware.
#[arg(long, value_enum)]
pub cache_prefix_provider: Option<CachePrefixProvider>,
// ---- session-affinity tuning ----
/// Header carrying the session ID for `--policy session_aware`.
@@ -106,6 +105,18 @@ pub struct Cli {
/// Session-affinity primary lookup and fallback behavior.
#[arg(long, value_enum)]
pub session_affinity_mode: Option<SessionAffinityMode>,
/// Disables the Session/Cache-Aware pressure guard.
#[arg(long)]
pub disable_pressure_guard: bool,
/// Absolute waiting-uncached-token gap required by the pressure guard.
#[arg(long)]
pub pressure_abs_threshold_tokens: Option<u64>,
/// Absolute millisecond gap when a Prefill queue estimate is available.
#[arg(long)]
pub pressure_abs_threshold_ms: Option<f64>,
/// Relative waiting-uncached-token multiplier required by the pressure guard.
#[arg(long)]
pub pressure_rel_threshold: Option<f64>,
/// Minimum cache-hit tokens for a cache-aware candidate.
#[arg(long)]
pub cache_affinity_min_matched_tokens: Option<u64>,
@@ -218,6 +229,11 @@ impl Cli {
/// (model id, static worker URLs).
pub fn into_config(self) -> Result<Config> {
let discovery = self.build_discovery()?;
let bucket_config = self
.bucket_config
.as_deref()
.map(load_bucket_config)
.transpose()?;
// Reject knobs that only take effect alongside another flag, rather
// than silently dropping them — mirrors the discovery mutual-exclusion
@@ -229,12 +245,16 @@ impl Cli {
enabled by --cb-threshold)"
));
}
let tuned_legacy_cache_aware = self.cache_threshold.is_some()
|| self.balance_abs_threshold.is_some()
|| self.balance_rel_threshold.is_some();
if tuned_legacy_cache_aware && self.policy != PolicyKind::CacheAwareZmq {
let cache_prefix_provider = self.cache_prefix_provider.unwrap_or_else(|| {
if self.kv_indexer_endpoint.is_some() {
CachePrefixProvider::Indexer
} else {
CachePrefixProvider::RadixTree
}
});
if self.cache_prefix_provider.is_some() && self.policy != PolicyKind::CacheAware {
return Err(anyhow!(
"cache-aware tuning flags require --policy cache_aware_zmq"
"--cache-prefix-provider requires --policy cache_aware"
));
}
if self.kv_indexer_query_timeout_ms == Some(0) {
@@ -257,22 +277,24 @@ impl Cli {
"--kv-indexer-query-max-inflight requires --kv-indexer-endpoint"
));
}
if self.kv_indexer_endpoint.is_some()
&& !matches!(
self.policy,
PolicyKind::CacheAware | PolicyKind::CacheAwareZmq
)
{
let cache_aware_uses_indexer = self.policy == PolicyKind::CacheAware
&& cache_prefix_provider == CachePrefixProvider::Indexer;
if self.kv_indexer_endpoint.is_some() && !cache_aware_uses_indexer {
if self.policy == PolicyKind::CacheAware {
return Err(anyhow!(
"--kv-indexer-endpoint requires --cache-prefix-provider indexer"
));
}
return Err(anyhow!(
"--kv-indexer-endpoint requires --policy cache_aware or cache_aware_zmq"
"--kv-indexer-endpoint requires --policy cache_aware"
));
}
if self.policy == PolicyKind::CacheAware && self.kv_indexer_endpoint.is_none() {
if cache_aware_uses_indexer && self.kv_indexer_endpoint.is_none() {
return Err(anyhow!(
"--policy cache_aware requires --kv-indexer-endpoint"
"--cache-prefix-provider indexer requires --kv-indexer-endpoint"
));
}
let tuned_cache_aware = tuned_legacy_cache_aware || self.kv_indexer_endpoint.is_some();
let tuned_cache_aware = self.policy == PolicyKind::CacheAware;
let affinity_policy = matches!(
self.policy,
PolicyKind::SessionAware | PolicyKind::CacheAware
@@ -289,6 +311,11 @@ impl Cli {
--session-affinity-mode require --policy session_aware"
));
}
if self.disable_pressure_guard && !affinity_policy {
return Err(anyhow!(
"--disable-pressure-guard requires --policy session_aware or cache_aware"
));
}
let tuned_cache_candidates = self.cache_affinity_min_matched_tokens.is_some()
|| self.cache_affinity_min_match_ratio.is_some()
|| self.cache_candidate_min_workers.is_some()
@@ -300,6 +327,15 @@ impl Cli {
"cache candidate tuning flags require --policy cache_aware"
));
}
if (self.pressure_abs_threshold_tokens.is_some()
|| self.pressure_abs_threshold_ms.is_some()
|| self.pressure_rel_threshold.is_some())
&& !affinity_policy
{
return Err(anyhow!(
"pressure guard tuning requires --policy session_aware or cache_aware"
));
}
let is_score_composition = matches!(
self.policy,
PolicyKind::FusedScore | PolicyKind::ScorePolicy
@@ -417,8 +453,26 @@ impl Cli {
let d = AffinityConfig::default();
let session_id_header = self.session_id_header.unwrap_or(d.session_id_header);
axum::http::HeaderName::try_from(session_id_header.as_str()).map_err(|e| {
anyhow!("--session-id-header {session_id_header:?} is not a valid HTTP header name: {e}")
anyhow!(
"--session-id-header {session_id_header:?} is not a valid HTTP header name: {e}"
)
})?;
let pressure_rel_threshold = self
.pressure_rel_threshold
.unwrap_or(d.pressure_rel_threshold);
if !pressure_rel_threshold.is_finite() || pressure_rel_threshold <= 1.0 {
return Err(anyhow!(
"--pressure-rel-threshold must be finite and greater than 1"
));
}
if self
.pressure_abs_threshold_ms
.is_some_and(|threshold| !threshold.is_finite() || threshold < 0.0)
{
return Err(anyhow!(
"--pressure-abs-threshold-ms must be finite and non-negative"
));
}
let cache_affinity_min_match_ratio = self
.cache_affinity_min_match_ratio
.or(d.cache_affinity_min_match_ratio);
@@ -473,6 +527,14 @@ impl Cli {
session_affinity_mode: self
.session_affinity_mode
.unwrap_or(d.session_affinity_mode),
pressure_guard: !self.disable_pressure_guard && d.pressure_guard,
pressure_abs_threshold_tokens: self
.pressure_abs_threshold_tokens
.unwrap_or(d.pressure_abs_threshold_tokens),
pressure_abs_threshold_ms: self
.pressure_abs_threshold_ms
.or(d.pressure_abs_threshold_ms),
pressure_rel_threshold,
cache_affinity_min_matched_tokens: self
.cache_affinity_min_matched_tokens
.or(d.cache_affinity_min_matched_tokens),
@@ -493,28 +555,22 @@ impl Cli {
cool_down_secs: self.cb_cool_down_secs.unwrap_or_else(default_cb_cool_down),
});
// Only build a CacheAwareConfig when the operator tuned at least
// one knob; otherwise leave it None so the policy uses its own
// defaults. Unset knobs fall back to the per-field defaults.
// Keep the selected prefix provider and optional Indexer settings
// together with the native Cache-Aware policy.
let kv_indexer_query_timeout_ms = self
.kv_indexer_query_timeout_ms
.unwrap_or(DEFAULT_KV_INDEXER_QUERY_TIMEOUT_MS);
let kv_indexer_query_max_inflight = self
.kv_indexer_query_max_inflight
.unwrap_or(DEFAULT_KV_INDEXER_QUERY_MAX_INFLIGHT);
let cache_aware = if tuned_cache_aware {
let d = CacheAwareConfig::default();
let kv_indexer_endpoint = self.kv_indexer_endpoint.map(|url| KvIndexerEndpointConfig {
url,
query_timeout_ms: self
.kv_indexer_query_timeout_ms
.unwrap_or(DEFAULT_KV_INDEXER_QUERY_TIMEOUT_MS),
query_max_inflight: self
.kv_indexer_query_max_inflight
.unwrap_or(DEFAULT_KV_INDEXER_QUERY_MAX_INFLIGHT),
query_timeout_ms: kv_indexer_query_timeout_ms,
query_max_inflight: kv_indexer_query_max_inflight,
});
Some(CacheAwareConfig {
cache_threshold: self.cache_threshold.unwrap_or(d.cache_threshold),
balance_abs_threshold: self
.balance_abs_threshold
.unwrap_or(d.balance_abs_threshold),
balance_rel_threshold: self
.balance_rel_threshold
.unwrap_or(d.balance_rel_threshold),
prefix_provider: cache_prefix_provider,
kv_indexer_endpoint,
})
} else {
@@ -536,6 +592,8 @@ impl Cli {
tokenizer_path: self.tokenizer_path.unwrap_or_else(|| self.model_id.clone()),
id: self.model_id,
policy: self.policy,
decode_policy: self.decode_policy,
bucket_config,
circuit_breaker,
cache_aware,
sticky,
@@ -570,13 +628,13 @@ impl Cli {
(true, true) => {
return Err(anyhow!(
"--worker-urls and --service-discovery are mutually exclusive; pass exactly one"
))
));
}
(false, false) => {
return Err(anyhow!(
"no discovery backend selected; pass --worker-urls <URL...> (static) \
or --service-discovery (kubernetes)"
))
));
}
(true, false) => {
if self.service_discovery_namespace.is_some()
@@ -613,6 +671,13 @@ impl Cli {
}
}
fn load_bucket_config(path: &str) -> Result<crate::config::BucketConfig> {
let raw = std::fs::read_to_string(path)
.map_err(|error| anyhow!("--bucket-config cannot read {path:?}: {error}"))?;
serde_json::from_str(&raw)
.map_err(|error| anyhow!("--bucket-config {path:?} is not valid JSON: {error}"))
}
/// Join space/repeated `key=value` selector terms into the single
/// comma-joined string the k8s backend's `labels_match_selector`
/// expects. `None` for an empty term list so [`resolve_mode`] can apply
@@ -930,6 +995,19 @@ mod tests {
);
}
#[test]
fn rejects_removed_cache_aware_zmq_policy() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"cache_aware_zmq",
]))
.unwrap_err()
.to_string();
assert!(err.contains("cache_aware_zmq"), "got: {err}");
}
#[test]
fn policy_accepts_only_routing_strategies() {
for value in ["prefix_cache", "overloaded"] {
@@ -1036,23 +1114,6 @@ mod tests {
);
}
#[test]
fn cache_aware_knob_builds_partial_config() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"cache_aware_zmq",
"--cache-threshold",
"0.7",
]))
.unwrap();
let ca = c.model.cache_aware.expect("cache_aware set");
assert_eq!(ca.cache_threshold, 0.7);
// Untouched knobs fall back to defaults.
assert_eq!(ca.balance_abs_threshold, 32);
}
#[test]
fn kv_indexer_reuses_cache_aware_policy_config() {
let c = into_config_owned(with_model(&[
@@ -1099,26 +1160,6 @@ mod tests {
assert_eq!(indexer.query_max_inflight, 32);
}
#[test]
fn kv_indexer_is_accepted_by_cache_aware_zmq() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"cache_aware_zmq",
"--kv-indexer-endpoint",
"http://indexer:50051",
]))
.unwrap();
let indexer = c
.model
.cache_aware
.expect("cache-aware config")
.kv_indexer_endpoint
.expect("Indexer config");
assert_eq!(indexer.url, "http://indexer:50051");
}
#[test]
fn kv_indexer_requires_cache_aware_policy() {
let err = into_config_owned(with_model(&[
@@ -1179,36 +1220,6 @@ mod tests {
assert!(err.contains("must be greater than zero"));
}
#[test]
fn no_cache_aware_flags_leaves_none() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"cache_aware_zmq",
]))
.unwrap();
assert!(c.model.cache_aware.is_none());
}
#[test]
fn rejects_cache_aware_knob_without_cache_aware_policy() {
// Default policy is round_robin, so a cache knob has no effect —
// reject rather than silently ignore it.
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--cache-threshold",
"0.7",
]))
.unwrap_err()
.to_string();
assert!(
err.contains("require --policy cache_aware_zmq"),
"got: {err}"
);
}
#[test]
fn log_format_parses_json() {
let c = into_config_owned(with_model(&[
@@ -1275,7 +1286,13 @@ mod tests {
for value in ["round_robin", "random", "power_of_two", "load_based"] {
assert!(choices.contains(value), "missing {value}: {choices}");
}
for value in ["fused_score", "cache_aware_zmq", "sticky"] {
for value in [
"fused_score",
"score_policy",
"session_aware",
"cache_aware",
"sticky",
] {
assert!(!choices.contains(value), "unexpected {value}: {choices}");
}
}
@@ -1424,14 +1441,14 @@ mod tests {
}
#[test]
fn rejects_cache_aware_zmq_as_sticky_fallback() {
fn rejects_cache_aware_as_sticky_fallback() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"sticky",
"--sticky-fallback-policy",
"cache_aware_zmq",
"cache_aware",
]))
.unwrap_err()
.to_string();
@@ -1630,17 +1647,21 @@ mod tests {
}
#[test]
fn rejects_removed_token_pressure_flags() {
for flag in [
"--disable-pressure-guard",
"--pressure-abs-threshold-tokens 2048",
"--pressure-rel-threshold 2.0",
] {
let error = cfg_of(&format!("--policy session_aware {flag}"))
.unwrap_err()
.to_string();
assert!(error.contains("unexpected argument"), "{flag}: {error}");
}
fn native_cache_pressure_flags_build_the_guard_contract() {
let config = cfg_of(
"--policy cache_aware --kv-indexer-endpoint http://indexer:50051 \
--disable-pressure-guard --pressure-abs-threshold-tokens 2048 \
--pressure-abs-threshold-ms 3.5 --pressure-rel-threshold 2.0",
)
.unwrap();
let affinity = config
.model
.affinity
.expect("cache-aware needs affinity config");
assert!(!affinity.pressure_guard);
assert_eq!(affinity.pressure_abs_threshold_tokens, 2_048);
assert_eq!(affinity.pressure_abs_threshold_ms, Some(3.5));
assert_eq!(affinity.pressure_rel_threshold, 2.0);
}
#[test]
@@ -1792,13 +1813,76 @@ mod tests {
}
#[test]
fn rejects_affinity_options_that_cannot_affect_the_selected_policy() {
let missing_indexer = cfg_of("--policy cache_aware")
.expect_err("cache_aware without an indexer can only behave like P2")
.to_string();
assert!(
missing_indexer.contains("--kv-indexer-endpoint"),
"got: {missing_indexer}"
fn cache_aware_defaults_to_router_radix_tree() {
let config = cfg_of("--policy cache_aware")
.expect("native cache-aware should not require an Indexer endpoint");
let cache = config
.model
.cache_aware
.expect("native cache-aware needs its default configuration");
assert_eq!(cache.prefix_provider, CachePrefixProvider::RadixTree);
assert!(cache.kv_indexer_endpoint.is_none());
}
#[test]
fn decode_policy_defaults_to_p2_and_accepts_legacy_compatibility_mode() {
let default_config = cfg_of("--policy power_of_two").unwrap();
assert_eq!(
default_config.model.decode_policy,
DecodePolicyKind::PowerOfTwo
);
let legacy_config = cfg_of("--decode-policy legacy_host_affinity").unwrap();
assert_eq!(
legacy_config.model.decode_policy,
DecodePolicyKind::LegacyHostAffinity
);
}
#[test]
fn bucket_config_json_is_loaded_and_validated_at_startup() {
let file = tempfile::NamedTempFile::new().unwrap();
std::fs::write(
file.path(),
r#"{
"ttft_slo_policy": "slo_first",
"tps_slo_policy": "best_effort",
"buckets": [
{
"id": "p-fast",
"stage": "prefill",
"rank": 10,
"worker_ids": ["http://worker:30000"],
"max_extend_tokens": 4096,
"max_context_tokens": 8192,
"ttft_p95_at_capacity_ms": 120
}
]
}"#,
)
.unwrap();
let path = file.path().to_str().unwrap().to_string();
let config = into_config_owned(with_model(&[
"--worker-urls",
"http://worker:30000",
"--bucket-config",
&path,
]))
.unwrap();
let buckets = config
.model
.bucket_config
.expect("Bucket config must be retained");
assert_eq!(buckets.buckets.len(), 1);
assert_eq!(buckets.buckets[0].id, "p-fast");
assert_eq!(
buckets.ttft_slo_policy,
crate::config::SloBucketPolicy::SloFirst
);
assert_eq!(
buckets.tps_slo_policy,
crate::config::SloBucketPolicy::BestEffort
);
}
}
+324
View File
@@ -15,6 +15,9 @@ impl Config {
if self.model.id.is_empty() {
return Err(anyhow!("model id must be non-empty"));
}
if let Some(bucket_config) = self.model.bucket_config.as_ref() {
validate_bucket_config(bucket_config)?;
}
match &self.discovery {
DiscoveryBackend::StaticUrls(s) => {
if s.urls.is_empty() {
@@ -66,6 +69,133 @@ impl Config {
}
}
fn validate_bucket_config(bucket_config: &BucketConfig) -> Result<()> {
if bucket_config.buckets.is_empty() {
return Err(anyhow!(
"bucket_config.buckets must be non-empty when configured"
));
}
let mut ids = std::collections::HashSet::new();
let mut ranks = std::collections::HashSet::new();
let mut stage_workers = std::collections::HashSet::new();
let mut has_prefill_bucket = false;
for bucket in &bucket_config.buckets {
has_prefill_bucket |= bucket.stage == BucketStage::Prefill;
if bucket.id.is_empty() || !ids.insert(bucket.id.as_str()) {
return Err(anyhow!(
"bucket_config bucket id must be non-empty and unique: {:?}",
bucket.id
));
}
if !ranks.insert((bucket.stage, bucket.rank)) {
return Err(anyhow!(
"bucket_config rank must be unique within each stage: {}",
bucket.rank
));
}
if bucket.worker_ids.is_empty() {
return Err(anyhow!(
"bucket_config bucket {:?} has no worker_ids",
bucket.id
));
}
let mut worker_ids = std::collections::HashSet::new();
for worker_id in &bucket.worker_ids {
if worker_id.is_empty() || !worker_ids.insert(worker_id.as_str()) {
return Err(anyhow!(
"bucket_config bucket {:?} has an empty or duplicate worker id",
bucket.id
));
}
if !stage_workers.insert((bucket.stage, worker_id.as_str())) {
return Err(anyhow!(
"bucket_config worker {:?} belongs to more than one {:?} bucket",
worker_id,
bucket.stage
));
}
}
validate_range(
bucket.min_extend_tokens,
bucket.max_extend_tokens,
&bucket.id,
"extend",
)?;
validate_range(
bucket.min_sequence_tokens,
bucket.max_sequence_tokens,
&bucket.id,
"sequence",
)?;
if bucket.max_context_tokens == Some(0) {
return Err(anyhow!(
"bucket_config bucket {:?} max_context_tokens must be > 0",
bucket.id
));
}
if bucket.ttft_p95_at_capacity_ms == Some(0) {
return Err(anyhow!(
"bucket_config bucket {:?} TTFT p95 must be > 0",
bucket.id
));
}
if bucket
.tps_p05_at_capacity
.is_some_and(|value| !value.is_finite() || value <= 0.0)
{
return Err(anyhow!(
"bucket_config bucket {:?} TPS p05 must be finite and > 0",
bucket.id
));
}
if bucket.max_pending_prefill_tokens == Some(0) {
return Err(anyhow!(
"bucket_config bucket {:?} max_pending_prefill_tokens must be > 0",
bucket.id
));
}
match bucket.stage {
BucketStage::Prefill
if bucket.min_sequence_tokens.is_some()
|| bucket.max_sequence_tokens.is_some()
|| bucket.tps_p05_at_capacity.is_some() =>
{
return Err(anyhow!(
"bucket_config Prefill bucket {:?} contains Decode-only sequence/TPS fields",
bucket.id
));
}
BucketStage::Decode
if bucket.min_extend_tokens.is_some()
|| bucket.max_extend_tokens.is_some()
|| bucket.ttft_p95_at_capacity_ms.is_some()
|| bucket.max_pending_prefill_tokens.is_some() =>
{
return Err(anyhow!(
"bucket_config Decode bucket {:?} contains Prefill-only extend/TTFT/pending fields",
bucket.id
));
}
_ => {}
}
}
if !has_prefill_bucket {
return Err(anyhow!(
"bucket_config must contain at least one Prefill bucket; enabling Bucket routing otherwise leaves every request without a Prefill domain"
));
}
Ok(())
}
fn validate_range(min: Option<u64>, max: Option<u64>, id: &str, name: &str) -> Result<()> {
if min.zip(max).is_some_and(|(min, max)| min > max) {
return Err(anyhow!(
"bucket_config bucket {id:?} has invalid {name} range: min > max"
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
@@ -85,6 +215,8 @@ mod tests {
id: model_id.into(),
tokenizer_path: "/tmp/tok.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: DecodePolicyKind::PowerOfTwo,
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -155,4 +287,196 @@ mod tests {
.to_string();
assert!(err.contains("unsupported scheme"), "got: {err}");
}
#[test]
fn rejects_a_worker_reused_by_two_buckets_of_the_same_stage() {
let mut config = cfg("qwen3", &["http://x:30000"]);
config.model.bucket_config = Some(BucketConfig {
buckets: vec![
BucketSpec {
id: "p-short".into(),
stage: BucketStage::Prefill,
rank: 10,
worker_ids: vec!["p1".into()],
min_extend_tokens: None,
max_extend_tokens: Some(1_024),
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: Some(100),
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
},
BucketSpec {
id: "p-long".into(),
stage: BucketStage::Prefill,
rank: 20,
worker_ids: vec!["p1".into()],
min_extend_tokens: Some(1_025),
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(8_192),
ttft_p95_at_capacity_ms: Some(200),
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
},
],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
});
let error = config.validate().unwrap_err().to_string();
assert!(error.contains("more than one"), "got: {error}");
}
#[test]
fn accepts_the_same_rank_in_independent_prefill_and_decode_stages() {
let mut config = cfg("qwen3", &["http://x:30000"]);
config.model.bucket_config = Some(BucketConfig {
buckets: vec![
BucketSpec {
id: "p-fast".into(),
stage: BucketStage::Prefill,
rank: 10,
worker_ids: vec!["p1".into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: Some(100),
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
},
BucketSpec {
id: "d-fast".into(),
stage: BucketStage::Decode,
rank: 10,
worker_ids: vec!["d1".into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: None,
tps_p05_at_capacity: Some(20.0),
max_pending_prefill_tokens: None,
},
],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::SloFirst,
});
config
.validate()
.expect("Prefill and Decode ranks only need to be unique within their stage");
}
#[test]
fn rejects_bucket_config_without_a_prefill_domain() {
let mut config = cfg("qwen3", &["http://x:30000"]);
config.model.bucket_config = Some(BucketConfig {
buckets: vec![BucketSpec {
id: "d-only".into(),
stage: BucketStage::Decode,
rank: 10,
worker_ids: vec!["d1".into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: None,
tps_p05_at_capacity: Some(20.0),
max_pending_prefill_tokens: None,
}],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::SloFirst,
});
let error = config.validate().unwrap_err().to_string();
assert!(error.contains("Prefill bucket"), "got: {error}");
}
#[test]
fn rejects_stage_inapplicable_bucket_fields_instead_of_ignoring_them() {
let mut prefill = cfg("qwen3", &["http://x:30000"]);
prefill.model.bucket_config = Some(BucketConfig {
buckets: vec![BucketSpec {
id: "p".into(),
stage: BucketStage::Prefill,
rank: 10,
worker_ids: vec!["p1".into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: Some(4_096),
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: Some(100),
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
}],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
});
let error = prefill.validate().unwrap_err().to_string();
assert!(error.contains("Decode-only"), "got: {error}");
let mut decode = cfg("qwen3", &["http://x:30000"]);
decode.model.bucket_config = Some(BucketConfig {
buckets: vec![
BucketSpec {
id: "p".into(),
stage: BucketStage::Prefill,
rank: 10,
worker_ids: vec!["p1".into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: Some(100),
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
},
BucketSpec {
id: "d".into(),
stage: BucketStage::Decode,
rank: 20,
worker_ids: vec!["d1".into()],
min_extend_tokens: Some(1),
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: Some(4_096),
max_context_tokens: Some(4_096),
ttft_p95_at_capacity_ms: None,
tps_p05_at_capacity: Some(20.0),
max_pending_prefill_tokens: None,
},
],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::SloFirst,
});
let error = decode.validate().unwrap_err().to_string();
assert!(error.contains("Prefill-only"), "got: {error}");
}
#[test]
fn bucket_json_rejects_unknown_profile_fields() {
let raw = r#"{
"buckets": [{
"id": "p-fast",
"stage": "prefill",
"rank": 10,
"worker_ids": ["p1"],
"ttft_p95_at_capcity_ms": 100
}]
}"#;
let error = serde_json::from_str::<BucketConfig>(raw)
.expect_err("a misspelled capacity profile must fail startup")
.to_string();
assert!(error.contains("ttft_p95_at_capcity_ms"), "got: {error}");
}
}
+103 -51
View File
@@ -1,3 +1,4 @@
use serde::Deserialize;
use std::num::NonZeroU32;
/// In-memory router configuration, built from CLI flags by
@@ -71,8 +72,7 @@ impl Default for ActiveLoadConfig {
///
/// Accepted on the CLI (`--policy`) as `round_robin` / `random` /
/// `power_of_two` / `load_based` / `fused_score` / `score_policy` /
/// `session_aware` / `cache_aware` / `cache_aware_zmq` /
/// `sticky`.
/// `session_aware` / `cache_aware` / `sticky`.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, clap::ValueEnum)]
pub enum PolicyKind {
#[default]
@@ -94,14 +94,9 @@ pub enum PolicyKind {
/// Selects a worker from session affinity.
#[value(name = "session_aware")]
SessionAware,
/// Selects cache-affine prefill candidates from external indexer data.
/// Selects cache-affine prefill candidates from the configured prefix provider.
#[value(name = "cache_aware")]
CacheAware,
/// Cache-aware routing fed by SGLang's ZMQ KV-cache event publisher.
/// Requires the model to have a tokenizer loaded; cache_aware tuning
/// lives on `ModelConfig::cache_aware`.
#[value(name = "cache_aware_zmq")]
CacheAwareZmq,
/// Sticky-session routing: pins a routing key (read from a
/// configurable request header) to a worker via an in-memory map, so
/// stateful sessions land on the same backend. Tuning — header name,
@@ -111,6 +106,71 @@ pub enum PolicyKind {
Sticky,
}
/// Policy used to select decode workers for PD requests.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, clap::ValueEnum)]
pub enum DecodePolicyKind {
#[default]
#[value(name = "power_of_two")]
PowerOfTwo,
#[value(name = "legacy_host_affinity")]
LegacyHostAffinity,
}
/// Role served by a static bucket.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BucketStage {
Prefill,
Decode,
}
/// SLO matching rules for a bucket.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SloBucketPolicy {
#[default]
Disabled,
BestEffort,
SloFirst,
}
/// Static bucket configuration loaded at Router startup.
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BucketConfig {
pub buckets: Vec<BucketSpec>,
#[serde(default)]
pub ttft_slo_policy: SloBucketPolicy,
#[serde(default)]
pub tps_slo_policy: SloBucketPolicy,
}
/// Runtime capacity assigned to one role. Lower ranks have higher priority.
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BucketSpec {
pub id: String,
pub stage: BucketStage,
pub rank: u32,
pub worker_ids: Vec<String>,
#[serde(default)]
pub min_extend_tokens: Option<u64>,
#[serde(default)]
pub max_extend_tokens: Option<u64>,
#[serde(default)]
pub min_sequence_tokens: Option<u64>,
#[serde(default)]
pub max_sequence_tokens: Option<u64>,
#[serde(default)]
pub max_context_tokens: Option<u64>,
#[serde(default)]
pub ttft_p95_at_capacity_ms: Option<u64>,
#[serde(default)]
pub tps_p05_at_capacity: Option<f64>,
#[serde(default)]
pub max_pending_prefill_tokens: Option<u64>,
}
impl std::fmt::Display for PolicyKind {
/// The CLI spelling for this policy kind.
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
@@ -231,8 +291,12 @@ pub struct ModelConfig {
/// is omitted. Resolved by [`crate::tokenizer::adapter::load`].
pub tokenizer_path: String,
pub policy: PolicyKind,
/// Selection policy for the decode pool.
pub decode_policy: DecodePolicyKind,
/// Optional static bucket configuration. `None` uses the global domain.
pub bucket_config: Option<BucketConfig>,
pub circuit_breaker: Option<CircuitBreakerConfig>,
/// Cache-Aware ZMQ tuning and optional external Indexer endpoint.
/// Cache-Aware prefix configuration.
pub cache_aware: Option<CacheAwareConfig>,
/// Tuning for the sticky-session policy. `Some` exactly when
/// `policy = "sticky"` (built by [`crate::config::cli::Cli::into_config`]).
@@ -306,48 +370,25 @@ fn parse_fuse_weight(name: &str, raw: &str) -> Result<f32, String> {
Ok(w)
}
/// Per-model cache-aware tuning.
#[derive(Debug, Clone)]
/// Cache-Aware prefix-match source.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, clap::ValueEnum)]
pub enum CachePrefixProvider {
#[default]
#[value(name = "radix_tree")]
RadixTree,
#[value(name = "indexer")]
Indexer,
}
/// Per-model Cache-Aware configuration.
#[derive(Debug, Clone, Default)]
pub struct CacheAwareConfig {
/// Lower bound on `matched_blocks / total_blocks` for the tree match
/// to win the selection. Below this, the policy falls back to
/// min-load. Default 0.5 — a half-cached prompt is still a strong
/// signal but not so weak that random hash collisions could trigger
/// affinity to an arbitrary worker.
pub cache_threshold: f32,
/// Absolute load spread (`max - min`) above which the cache check is
/// skipped in favour of min-load. Default 32 — picked to dominate
/// over typical batch-of-8 effect.
pub balance_abs_threshold: usize,
/// Multiplicative load spread (`max > min * balance_rel_threshold`)
/// that the absolute check is gated on. Default 1.1 — 10 % relative
/// difference triggers re-balancing.
pub balance_rel_threshold: f32,
/// Optional external KV Indexer client configuration.
/// Prefix-match source for native Cache-Aware.
pub prefix_provider: CachePrefixProvider,
/// External Indexer configuration when `prefix_provider = indexer`.
pub kv_indexer_endpoint: Option<KvIndexerEndpointConfig>,
}
impl Default for CacheAwareConfig {
fn default() -> Self {
Self {
cache_threshold: default_cache_threshold(),
balance_abs_threshold: default_balance_abs(),
balance_rel_threshold: default_balance_rel(),
kv_indexer_endpoint: None,
}
}
}
fn default_cache_threshold() -> f32 {
0.5
}
fn default_balance_abs() -> usize {
32
}
fn default_balance_rel() -> f32 {
1.1
}
/// Default routing-key header for the sticky policy. The `x-sgl-` prefix
/// matches the router's other emitted/consumed metadata headers
/// (`x-sgl-decode-url`, `x-sgl-router-error-code`).
@@ -395,6 +436,10 @@ pub struct AffinityConfig {
pub stable_pair: bool,
pub mode: AffinityMode,
pub session_affinity_mode: SessionAffinityMode,
pub pressure_guard: bool,
pub pressure_abs_threshold_tokens: u64,
pub pressure_abs_threshold_ms: Option<f64>,
pub pressure_rel_threshold: f64,
pub cache_affinity_min_matched_tokens: Option<u64>,
pub cache_affinity_min_match_ratio: Option<f64>,
pub cache_candidate_min_workers: usize,
@@ -412,6 +457,10 @@ impl Default for AffinityConfig {
stable_pair: false,
mode: AffinityMode::Soft,
session_affinity_mode: SessionAffinityMode::Bucket,
pressure_guard: true,
pressure_abs_threshold_tokens: 1_024,
pressure_abs_threshold_ms: None,
pressure_rel_threshold: 1.5,
// Indexer prefix scans are truncated, so use an absolute token floor.
cache_affinity_min_matched_tokens: Some(1_024),
cache_affinity_min_match_ratio: None,
@@ -434,8 +483,7 @@ pub struct StickyConfig {
/// Policy used to pick a worker when a request has no routing key, and
/// to pick the initial worker when a new key is first seen. One of
/// `round_robin` / `random` / `power_of_two` / `load_based` — the
/// dependency-free policies the factory can build standalone (no
/// `HashTree` / tokenizer / ZMQ feed).
/// dependency-free policies the factory can build standalone.
pub fallback_policy: StickyFallbackKind,
/// Evict an assignment after it has been idle (unreferenced) this many
/// seconds. Bounds the map against unbounded routing-key cardinality.
@@ -553,9 +601,13 @@ pub enum K8sDiscoveryMode {
/// invalid.
#[derive(Debug, thiserror::Error)]
pub enum ConfigError {
#[error("discovery.k8s requires either `label_selector` (plain) or both `prefill_selector` and `decode_selector` (PD); none were set")]
#[error(
"discovery.k8s requires either `label_selector` (plain) or both `prefill_selector` and `decode_selector` (PD); none were set"
)]
NoSelector,
#[error("discovery.k8s: `label_selector` (plain) and `prefill_selector`/`decode_selector` (PD) are mutually exclusive — set one or the other, not both")]
#[error(
"discovery.k8s: `label_selector` (plain) and `prefill_selector`/`decode_selector` (PD) are mutually exclusive — set one or the other, not both"
)]
MixedModes,
#[error("discovery.k8s: PD mode requires BOTH `prefill_selector` and `decode_selector`")]
PartialPdSelectors,
+52 -17
View File
@@ -3,7 +3,7 @@
use anyhow::{Context, Result};
use clap::Parser;
use sgl_router::config::{Cli, LogFormat};
use sgl_router::config::{CachePrefixProvider, Cli, LogFormat, PolicyKind};
use std::sync::Arc;
use tokio::signal::unix::{signal, Signal, SignalKind};
@@ -86,6 +86,7 @@ async fn main() -> Result<()> {
init_tracing(&cfg.observability.log_level, cfg.observability.log_format)?;
tracing::info!(
configured_decode_policy = ?cfg.model.decode_policy,
"sgl-router {} starting on {}:{}",
env!("CARGO_PKG_VERSION"),
cfg.server.host,
@@ -98,27 +99,27 @@ async fn main() -> Result<()> {
);
let registry = Arc::new(sgl_router::workers::WorkerRegistry::default());
let prefix_index = cfg
.model
.cache_aware
.as_ref()
let cache_aware_uses_indexer = cfg.model.policy == PolicyKind::CacheAware
&& cfg
.model
.cache_aware
.as_ref()
.is_some_and(|cache| cache.prefix_provider == CachePrefixProvider::Indexer);
let prefix_index: Option<Arc<dyn sgl_kv_indexer::PrefixIndex>> = cache_aware_uses_indexer
.then_some(cfg.model.cache_aware.as_ref())
.flatten()
.and_then(|cache| cache.kv_indexer_endpoint.as_ref())
.map(|indexer| {
let config = sgl_kv_indexer::PrefixIndexConfig {
endpoint: indexer.url.clone(),
query_deadline: std::time::Duration::from_millis(indexer.query_timeout_ms),
max_inflight: indexer.query_max_inflight,
};
let config = prefix_index_config(indexer);
sgl_kv_indexer::GrpcPrefixIndex::new(config)
.map(Arc::new)
.map(|index| Arc::new(index) as Arc<dyn sgl_kv_indexer::PrefixIndex>)
.context("configure KV Indexer client")
})
.transpose()?;
// Build the KV-event index up front so the cache-aware-zmq policy can
// share its `HashTree` handle + `BlockSizeOracle`. An external Indexer makes
// the local tree irrelevant to routing, so only discover hash metadata rather
// than duplicating every KV event.
// Build the local prefix index and block metadata used by the Radix Tree
// provider. An external Indexer only needs hash metadata, so it does not
// subscribe to the local KV-event stream.
let block_size_oracle = sgl_router::policies::kv_events::BlockSizeOracle::new();
let kv_event_http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(2))
@@ -139,9 +140,7 @@ async fn main() -> Result<()> {
sgl_router::policies::factory::build_registry(
&cfg,
kv_index.tree(),
Arc::clone(&tokenizers),
Arc::clone(&block_size_oracle),
kv_index.engine_load(),
)
.context("build policy registry")?,
);
@@ -197,6 +196,18 @@ async fn main() -> Result<()> {
active_load,
);
app_ctx.prefix_index = prefix_index;
app_ctx.radix_tree_prefix_provider = (cfg.model.policy == PolicyKind::CacheAware
&& cfg
.model
.cache_aware
.as_ref()
.is_some_and(|cache| cache.prefix_provider == CachePrefixProvider::RadixTree))
.then(|| {
sgl_router::policies::prefix_provider::RadixTreePrefixProvider::new(
kv_index.tree(),
Arc::clone(&block_size_oracle),
)
});
app_ctx.block_size_oracle = block_size_oracle;
app_ctx.engine_load = kv_index.engine_load();
let ctx = Arc::new(app_ctx);
@@ -225,6 +236,18 @@ async fn main() -> Result<()> {
server_result
}
/// Build the external Indexer client with the Router's bounded query settings.
fn prefix_index_config(
indexer: &sgl_router::config::KvIndexerEndpointConfig,
) -> sgl_kv_indexer::PrefixIndexConfig {
sgl_kv_indexer::PrefixIndexConfig {
endpoint: indexer.url.clone(),
query_deadline: std::time::Duration::from_millis(indexer.query_timeout_ms),
max_inflight: indexer.query_max_inflight,
}
}
/// Waits for either Unix termination signal and logs the selected cause.
async fn shutdown_signal(mut sigterm: Signal, mut sigint: Signal) {
tokio::select! {
_ = sigterm.recv() => tracing::info!("got SIGTERM, shutting down"),
@@ -236,6 +259,18 @@ async fn shutdown_signal(mut sigterm: Signal, mut sigint: Signal) {
mod tests {
use super::*;
#[test]
fn prefix_index_config_preserves_router_limits() {
let config = prefix_index_config(&sgl_router::config::KvIndexerEndpointConfig {
url: "http://127.0.0.1:50051".to_string(),
query_timeout_ms: 25,
query_max_inflight: 17,
});
assert_eq!(config.endpoint, "http://127.0.0.1:50051");
assert_eq!(config.query_deadline, std::time::Duration::from_millis(25));
assert_eq!(config.max_inflight, 17);
}
#[tokio::test]
async fn install_signal_handlers_returns_both() {
// Pins the contract that handler installation works on a standard
@@ -4,11 +4,8 @@
//! Per-worker active-load tracking with RAII guards and a stale-request
//! janitor.
//!
//! The cache-aware-zmq policy ([`super::cache_aware_zmq`]) needs to combine
//! the hash tree's overlap score with a per-worker load signal. The
//! per-worker `Worker::active_requests` counter tracks one axis — number of
//! in-flight HTTP requests — and is already drop-safe through
//! [`crate::workers::LoadGuard`].
//! The per-worker `Worker::active_requests` counter tracks in-flight HTTP
//! requests and is drop-safe through [`crate::workers::LoadGuard`].
//!
//! This module adds two things on top of that:
//!
@@ -72,10 +69,7 @@ impl std::fmt::Display for RequestId {
}
}
/// Per-worker counters: one for prefill (token) load, one for decode (block)
/// load. The two axes are tracked separately so cache-aware-zmq can score
/// prefill candidates by token load and decode candidates by block load
/// without each axis spamming through the other's counter.
/// Per-worker counters for prefill and decode work.
///
/// Production tracks **active requests** as the unit (count of in-flight
/// requests pinning the worker), not raw token / block counts — until the
@@ -172,10 +166,8 @@ impl Clock for MockClock {
/// Registry of in-flight requests + per-worker active-load counters.
///
/// Constructed once per `AppContext`; the cache-aware-zmq policy reads
/// per-worker `prefill_load` / `decode_load` from here when scoring
/// candidates, and the proxy holds an [`ActiveLoadGuard`] per request so
/// counters decrement on drop. A background task periodically calls
/// Constructed once per `AppContext`; the proxy holds an [`ActiveLoadGuard`]
/// per request so counters decrement on drop. A background task periodically calls
/// [`Self::sweep_stale`] to evict requests that outlived
/// `stale_request_timeout`.
#[derive(Debug)]
+389 -71
View File
@@ -1,19 +1,24 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Shared admission and candidate comparison for Prefill and Decode.
//! Shared capacity admission and pressure guards for prefill and decode.
//!
//! Decisions use only fields published in `LoadStat`.
//! Native Cache-Aware uses monitor-backed admission only when every expected
//! DP rank has a fresh, complete #34608 ZMQ sample. Otherwise it falls back to
//! Router-local load.
use crate::policies::engine_load::{EngineLoadSnapshot, EngineWorkerLoad};
use crate::policies::engine_load::{EngineLoadSnapshot, NativeCacheWorkerLoad};
use crate::policies::power_of_two::select_with_snapshot;
use crate::policies::{CacheCandidate, CacheCandidateProposal, SelectionProposal};
use crate::policies::{CacheCandidate, CacheCandidateProposal, GuardHints, SelectionProposal};
use crate::workers::Worker;
use std::cmp::Ordering;
use std::collections::HashMap;
use std::sync::Arc;
/// A Prefill candidate range and its optional pending-token budget.
/// A prefill candidate domain and its optional queue budget.
///
/// `max_pending_prefill_tokens` is enforced only when the native monitor
/// provides `num_waiting_uncached_tokens`.
pub struct CandidateRange<'a> {
pub id: &'a str,
pub workers: &'a [Arc<Worker>],
@@ -30,7 +35,7 @@ impl<'a> CandidateRange<'a> {
}
}
/// A role-specific candidate domain resolved before policy evaluation.
/// Role-specific candidate domains resolved before policy selection.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RoutingStage {
Prefill,
@@ -100,8 +105,7 @@ pub enum DecisionReason {
Primary,
CacheCandidate,
BackupPrimaryAdmission,
/// Both Decode candidates were admitted; the lower-pressure backup won.
BackupLoadComparison,
BackupPressureGuard,
RangeFallback,
CapacityFallbackPowerOfTwo,
}
@@ -116,12 +120,22 @@ pub struct FinalDecision {
pub load_snapshot_version: u64,
}
/// Resolves a bounded cache-candidate set, using pressure to break near ties.
/// Cache-Aware selection audit data. These fields do not affect selection.
pub struct CacheCandidateResolution {
pub decision: Option<FinalDecision>,
pub prefill_pressure_source: &'static str,
pub admission_evaluated_candidates: u64,
pub admission_rejected_candidates: u64,
pub pressure_guard_compared_pairs: u64,
pub pressure_guard_overrides: u64,
}
/// Selects a worker from bounded cache candidates and records guard coverage.
pub fn resolve_cache_candidates(
proposal: &CacheCandidateProposal,
request_input_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> Option<FinalDecision> {
) -> CacheCandidateResolution {
let loads = FreshLoadLookup::new(
Some(snapshot),
proposal
@@ -134,29 +148,64 @@ pub fn resolve_cache_candidates(
.iter()
.filter(|candidate| is_cache_candidate_admitted(candidate, request_input_tokens, &loads))
.collect();
let work_floor = admitted
let admission_rejected_candidates =
proposal.candidates.len().saturating_sub(admitted.len()) as u64;
let Some(work_floor) = admitted
.iter()
.copied()
.min_by_key(|candidate| candidate.uncached_tokens)?;
.min_by_key(|candidate| candidate.uncached_tokens)
else {
return CacheCandidateResolution {
decision: None,
prefill_pressure_source: loads.prefill_pressure_source(),
admission_evaluated_candidates: proposal.candidates.len() as u64,
admission_rejected_candidates,
pressure_guard_compared_pairs: 0,
pressure_guard_overrides: 0,
};
};
let near_tie_ceiling = work_floor
.uncached_tokens
.saturating_add(proposal.cache_switch_margin_tokens);
let mut winner = work_floor;
let mut pressure_guard_compared_pairs = 0;
let mut pressure_guard_overrides = 0;
for candidate in admitted {
if candidate.uncached_tokens <= near_tie_ceiling
&& compare_cache_candidates(winner, candidate, &loads).is_gt()
if candidate.worker.id == winner.worker.id || candidate.uncached_tokens > near_tie_ceiling {
continue;
}
let baseline = compare_cache_candidates(winner, candidate, proposal, &loads, false);
let ordering = if proposal.enable_pressure_guard
&& cache_pressure_guard_comparable(winner, candidate, &loads)
{
pressure_guard_compared_pairs += 1;
let guarded = compare_cache_candidates(winner, candidate, proposal, &loads, true);
if guarded != baseline {
pressure_guard_overrides += 1;
}
guarded
} else {
baseline
};
if ordering.is_gt() {
winner = candidate;
}
}
Some(FinalDecision {
selected: Arc::clone(&winner.worker),
primary: Arc::clone(&winner.worker),
backup: None,
reason: DecisionReason::CacheCandidate,
candidate_range_id: winner.candidate_range_id.clone(),
load_snapshot_version: snapshot.version,
})
CacheCandidateResolution {
decision: Some(FinalDecision {
selected: Arc::clone(&winner.worker),
primary: Arc::clone(&winner.worker),
backup: None,
reason: DecisionReason::CacheCandidate,
candidate_range_id: winner.candidate_range_id.clone(),
load_snapshot_version: snapshot.version,
}),
prefill_pressure_source: loads.prefill_pressure_source(),
admission_evaluated_candidates: proposal.candidates.len() as u64,
admission_rejected_candidates,
pressure_guard_compared_pairs,
pressure_guard_overrides,
}
}
pub fn resolve_prefill(
@@ -164,6 +213,35 @@ pub fn resolve_prefill(
proposal: &SelectionProposal,
request_input_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> Option<FinalDecision> {
resolve_prefill_admitted(range, proposal, request_input_tokens, snapshot).or_else(|| {
if !contains_worker(range, &proposal.primary) {
return None;
}
let backup = proposal
.backup
.as_ref()
.filter(|worker| contains_worker(range, worker))
.cloned();
let legal = legal_prefill_candidates(range, proposal);
let selected = select_with_snapshot(&legal, Some(snapshot))?;
Some(FinalDecision {
selected,
primary: Arc::clone(&proposal.primary),
backup,
reason: DecisionReason::CapacityFallbackPowerOfTwo,
candidate_range_id: range.id.to_string(),
load_snapshot_version: snapshot.version,
})
})
}
/// Resolves prefill admission without overcommitting a full candidate range.
pub fn resolve_prefill_admitted(
range: &CandidateRange<'_>,
proposal: &SelectionProposal,
request_input_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> Option<FinalDecision> {
if !contains_worker(range, &proposal.primary) {
return None;
@@ -174,21 +252,30 @@ pub fn resolve_prefill(
.filter(|worker| contains_worker(range, worker))
.cloned();
let primary_admitted = is_proposal_worker_eligible(proposal, &proposal.primary)
&& is_prefill_admitted(&proposal.primary, request_input_tokens, snapshot);
&& is_prefill_admitted(range, &proposal.primary, request_input_tokens, snapshot);
let backup_admitted = backup.as_ref().is_some_and(|worker| {
is_proposal_worker_eligible(proposal, worker)
&& is_prefill_admitted(worker, request_input_tokens, snapshot)
&& is_prefill_admitted(range, worker, request_input_tokens, snapshot)
});
let (selected, reason) = match (primary_admitted, backup.as_ref(), backup_admitted) {
(true, Some(backup), true) => {
if pressure_guard_prefers_backup(
&proposal.primary,
backup,
&proposal.guard_hints,
snapshot,
) {
(Arc::clone(backup), DecisionReason::BackupPressureGuard)
} else {
(Arc::clone(&proposal.primary), DecisionReason::Primary)
}
}
(true, _, _) => (Arc::clone(&proposal.primary), DecisionReason::Primary),
(false, Some(backup), true) => (Arc::clone(backup), DecisionReason::BackupPrimaryAdmission),
_ => {
let legal = legal_prefill_candidates(range, proposal);
range_fallback(&legal, request_input_tokens, snapshot).or_else(|| {
select_with_snapshot(&legal, Some(snapshot))
.map(|worker| (worker, DecisionReason::CapacityFallbackPowerOfTwo))
})?
range_fallback(range, &legal, request_input_tokens, snapshot)?
}
};
Some(FinalDecision {
@@ -222,7 +309,7 @@ pub fn resolve_decode(
let (selected, reason) = match (primary_admitted, backup.as_ref(), backup_admitted) {
(true, Some(backup), true) => {
if compare_decode_pressure(&proposal.primary, backup, Some(snapshot)).is_gt() {
(Arc::clone(backup), DecisionReason::BackupLoadComparison)
(Arc::clone(backup), DecisionReason::BackupPressureGuard)
} else {
(Arc::clone(&proposal.primary), DecisionReason::Primary)
}
@@ -259,24 +346,31 @@ fn is_proposal_worker_eligible(proposal: &SelectionProposal, candidate: &Arc<Wor
.is_none_or(|workers| workers.iter().any(|worker| worker.id == candidate.id))
}
/// A zero LoadStat capacity is unknown and does not reject a candidate.
fn has_kv_capacity(load: Option<&EngineWorkerLoad>, requested_tokens: u64) -> bool {
/// Applies snapshot-backed capacity admission when native monitor data is complete.
/// Workers without monitor data remain eligible and use Router-local ordering.
fn has_kv_capacity(load: Option<&NativeCacheWorkerLoad>, requested_tokens: u64) -> bool {
let Some(load) = load else {
return true;
};
load.max_total_num_tokens == 0
|| load.num_tokens.saturating_add(requested_tokens) <= load.max_total_num_tokens
load.num_running_reqs.saturating_add(1) <= load.max_running_requests
&& load.num_total_tokens.saturating_add(requested_tokens) <= load.max_total_num_tokens
}
fn is_prefill_admitted(
range: &CandidateRange<'_>,
worker: &Arc<Worker>,
request_input_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> bool {
has_kv_capacity(
snapshot.fresh_load_for_url(&worker.url),
request_input_tokens,
)
let load = snapshot.fresh_native_cache_load_for_url(&worker.url);
has_kv_capacity(load, request_input_tokens)
&& range.max_pending_prefill_tokens.is_none_or(|limit| {
load.is_none_or(|load| {
load.num_waiting_uncached_tokens
.saturating_add(request_input_tokens)
<= limit
})
})
}
fn is_decode_admitted(
@@ -284,7 +378,10 @@ fn is_decode_admitted(
request_kv_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> bool {
has_kv_capacity(snapshot.fresh_load_for_url(&worker.url), request_kv_tokens)
has_kv_capacity(
snapshot.fresh_native_cache_load_for_url(&worker.url),
request_kv_tokens,
)
}
fn is_cache_candidate_admitted(
@@ -292,25 +389,114 @@ fn is_cache_candidate_admitted(
request_input_tokens: u64,
loads: &FreshLoadLookup<'_>,
) -> bool {
has_kv_capacity(loads.get(&candidate.worker.id), request_input_tokens)
let Some(load) = loads.get(&candidate.worker.id) else {
return true;
};
has_kv_capacity(Some(load), request_input_tokens)
&& candidate.max_pending_prefill_tokens.is_none_or(|limit| {
load.num_waiting_uncached_tokens
.saturating_add(candidate.uncached_tokens)
<= limit
})
}
fn compare_cache_candidates(
left: &CacheCandidate,
right: &CacheCandidate,
proposal: &CacheCandidateProposal,
loads: &FreshLoadLookup<'_>,
enable_pressure_guard: bool,
) -> Ordering {
let work_delta = left.uncached_tokens.abs_diff(right.uncached_tokens);
if work_delta > proposal.cache_switch_margin_tokens {
return left
.uncached_tokens
.cmp(&right.uncached_tokens)
.then_with(|| loads.compare_prefill_pressure(&left.worker, &right.worker))
.then_with(|| left.worker.id.0.cmp(&right.worker.id.0));
}
if enable_pressure_guard {
if materially_more_pressured(
&left.worker,
&right.worker,
proposal.pressure_abs_threshold_tokens,
proposal.pressure_abs_threshold_ms,
proposal.pressure_rel_threshold,
loads,
) {
return Ordering::Greater;
}
if materially_more_pressured(
&right.worker,
&left.worker,
proposal.pressure_abs_threshold_tokens,
proposal.pressure_abs_threshold_ms,
proposal.pressure_rel_threshold,
loads,
) {
return Ordering::Less;
}
}
left.uncached_tokens
.cmp(&right.uncached_tokens)
.then_with(|| loads.compare_prefill_pressure(&left.worker, &right.worker))
.then_with(|| left.worker.id.0.cmp(&right.worker.id.0))
}
/// Per-request lookup that uses engine pressure only for a complete fresh set.
fn cache_pressure_guard_comparable(
left: &CacheCandidate,
right: &CacheCandidate,
loads: &FreshLoadLookup<'_>,
) -> bool {
loads.comparable_get(&left.worker.id).is_some()
&& loads.comparable_get(&right.worker.id).is_some()
}
fn materially_more_pressured(
candidate: &Arc<Worker>,
other: &Arc<Worker>,
absolute_threshold_tokens: u64,
absolute_threshold_ms: Option<f64>,
relative_threshold: f64,
loads: &FreshLoadLookup<'_>,
) -> bool {
let (Some(candidate_load), Some(other_load)) = (
loads.comparable_get(&candidate.id),
loads.comparable_get(&other.id),
) else {
return false;
};
if let Some(absolute_threshold_ms) = absolute_threshold_ms.filter(|_| {
candidate_load.estimated_prefill_queue_ms.is_some()
&& other_load.estimated_prefill_queue_ms.is_some()
}) {
let candidate_pressure = candidate_load
.estimated_prefill_queue_ms
.expect("availability was checked");
let other_pressure = other_load
.estimated_prefill_queue_ms
.expect("availability was checked");
return candidate_pressure - other_pressure > absolute_threshold_ms
&& candidate_pressure > other_pressure * relative_threshold;
}
candidate_load
.num_waiting_uncached_tokens
.saturating_sub(other_load.num_waiting_uncached_tokens)
> absolute_threshold_tokens
&& candidate_load.num_waiting_uncached_tokens as f64
> other_load.num_waiting_uncached_tokens as f64 * relative_threshold
}
/// Constant-time request view over one captured load snapshot.
///
/// External values are compared only when every candidate is present. Mixed
/// candidate sets use Router-local active load to preserve ordering.
pub(crate) struct FreshLoadLookup<'a> {
by_worker_id: HashMap<String, &'a EngineWorkerLoad>,
by_worker_id: HashMap<String, &'a NativeCacheWorkerLoad>,
basic_by_worker_id: HashMap<String, &'a crate::policies::engine_load::EngineWorkerLoad>,
local_active_by_worker_id: HashMap<String, usize>,
compare_engine: bool,
compare_basic_engine: bool,
}
impl<'a> FreshLoadLookup<'a> {
@@ -324,6 +510,16 @@ impl<'a> FreshLoadLookup<'a> {
.map(|worker| (worker.id.0.clone(), worker.active_load()))
.collect();
let by_worker_id = snapshot
.into_iter()
.flat_map(|snapshot| {
workers.iter().filter_map(move |worker| {
snapshot
.fresh_native_cache_load_for_url(&worker.url)
.map(|load| (worker.id.0.clone(), load))
})
})
.collect::<HashMap<_, _>>();
let basic_by_worker_id = snapshot
.into_iter()
.flat_map(|snapshot| {
workers.iter().filter_map(move |worker| {
@@ -335,24 +531,28 @@ impl<'a> FreshLoadLookup<'a> {
.collect::<HashMap<_, _>>();
let compare_engine = !local_active_by_worker_id.is_empty()
&& by_worker_id.len() == local_active_by_worker_id.len();
let compare_basic_engine = !local_active_by_worker_id.is_empty()
&& basic_by_worker_id.len() == local_active_by_worker_id.len();
Self {
by_worker_id,
basic_by_worker_id,
local_active_by_worker_id,
compare_engine,
compare_basic_engine,
}
}
pub(crate) fn get(
&self,
worker_id: &crate::discovery::WorkerId,
) -> Option<&'a EngineWorkerLoad> {
) -> Option<&'a NativeCacheWorkerLoad> {
self.by_worker_id.get(worker_id.0.as_str()).copied()
}
fn comparable_get(
&self,
worker_id: &crate::discovery::WorkerId,
) -> Option<&'a EngineWorkerLoad> {
) -> Option<&'a NativeCacheWorkerLoad> {
self.compare_engine.then(|| self.get(worker_id)).flatten()
}
@@ -369,8 +569,7 @@ impl<'a> FreshLoadLookup<'a> {
fn compare_prefill_keys(&self, left: &PressureKey<'a>, right: &PressureKey<'a>) -> Ordering {
match (left.load, right.load) {
(Some(left_load), Some(right_load)) => prefill_pressure_key(left_load)
.cmp(&prefill_pressure_key(right_load))
(Some(left_load), Some(right_load)) => compare_prefill_load(left_load, right_load)
.then_with(|| left.local_active.cmp(&right.local_active)),
_ => left.local_active.cmp(&right.local_active),
}
@@ -392,10 +591,30 @@ impl<'a> FreshLoadLookup<'a> {
self.compare_prefill_keys(&self.pressure_key(left), &self.pressure_key(right))
}
/// Returns corrected engine queue depth for a complete fresh set, otherwise
/// local load.
pub(crate) fn prefill_pressure_source(&self) -> &'static str {
if self.compare_engine
&& self
.by_worker_id
.values()
.all(|load| load.estimated_prefill_queue_ms.is_some())
{
"estimated_prefill_queue_ms"
} else if self.compare_engine {
"native_queue_tokens"
} else {
"router_local"
}
}
/// Returns a queue depth consistent with admission for this request.
///
/// A fully covered candidate set uses `waiting + running`; otherwise the
/// whole set uses Router-local active load. Dispatches after the snapshot
/// are added to the reported value.
pub(crate) fn score_load(&self, worker: &Arc<Worker>) -> usize {
self.comparable_get(&worker.id)
self.compare_basic_engine
.then(|| self.basic_by_worker_id.get(worker.id.0.as_str()).copied())
.flatten()
.map(|load| {
let recent_dispatches = worker
.slots_acquired_since(load.captured_at)
@@ -414,7 +633,6 @@ impl<'a> FreshLoadLookup<'a> {
.unwrap_or(usize::MAX)
})
}
fn min_by_pressure_key(
&self,
candidates: Vec<Arc<Worker>>,
@@ -435,18 +653,20 @@ impl<'a> FreshLoadLookup<'a> {
}
struct PressureKey<'a> {
load: Option<&'a EngineWorkerLoad>,
load: Option<&'a NativeCacheWorkerLoad>,
local_active: usize,
}
fn range_fallback(
range: &CandidateRange<'_>,
legal: &[Arc<Worker>],
request_input_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> Option<(Arc<Worker>, DecisionReason)> {
let admitted = legal
.iter()
.filter(|worker| is_prefill_admitted(worker, request_input_tokens, snapshot))
.filter(|worker| contains_worker(range, worker))
.filter(|worker| is_prefill_admitted(range, worker, request_input_tokens, snapshot))
.cloned()
.collect::<Vec<_>>();
let loads = FreshLoadLookup::new(Some(snapshot), admitted.iter());
@@ -486,7 +706,7 @@ fn decode_domain_fallback(
.map(|worker| (worker, DecisionReason::RangeFallback))
}
/// Compares Prefill pressure by waiting requests, running requests, and KV use.
/// Compares prefill pressure by queue time when available, then by the V3 load tuple.
pub(crate) fn compare_prefill_pressure(
left: &Arc<Worker>,
right: &Arc<Worker>,
@@ -494,27 +714,37 @@ pub(crate) fn compare_prefill_pressure(
) -> Ordering {
match snapshot.and_then(|snapshot| {
Some((
snapshot.fresh_load_for_url(&left.url)?,
snapshot.fresh_load_for_url(&right.url)?,
snapshot.fresh_native_cache_load_for_url(&left.url)?,
snapshot.fresh_native_cache_load_for_url(&right.url)?,
))
}) {
Some((left_load, right_load)) => prefill_pressure_key(left_load)
.cmp(&prefill_pressure_key(right_load))
Some((left_load, right_load)) => compare_prefill_load(left_load, right_load)
.then_with(|| left.active_load().cmp(&right.active_load())),
None => left.active_load().cmp(&right.active_load()),
}
}
fn prefill_pressure_key(load: &EngineWorkerLoad) -> (u64, u64, u64, u64) {
fn prefill_pressure_key(load: &NativeCacheWorkerLoad) -> (u64, u64, u64) {
(
load.num_waiting_uncached_tokens,
load.num_waiting_reqs,
load.num_running_reqs,
load.num_tokens,
load.max_total_num_tokens,
)
}
/// Compares Decode pressure using only LoadStat values.
fn compare_prefill_load(left: &NativeCacheWorkerLoad, right: &NativeCacheWorkerLoad) -> Ordering {
match (
left.estimated_prefill_queue_ms,
right.estimated_prefill_queue_ms,
) {
(Some(left_ms), Some(right_ms)) => left_ms
.total_cmp(&right_ms)
.then_with(|| prefill_pressure_key(left).cmp(&prefill_pressure_key(right))),
_ => prefill_pressure_key(left).cmp(&prefill_pressure_key(right)),
}
}
/// Compares decode pressure from LoadStat without treating unknown capacity as zero.
pub(crate) fn compare_decode_pressure(
left: &Arc<Worker>,
right: &Arc<Worker>,
@@ -522,8 +752,8 @@ pub(crate) fn compare_decode_pressure(
) -> Ordering {
match snapshot.and_then(|snapshot| {
Some((
snapshot.fresh_load_for_url(&left.url)?,
snapshot.fresh_load_for_url(&right.url)?,
snapshot.fresh_native_cache_load_for_url(&left.url)?,
snapshot.fresh_native_cache_load_for_url(&right.url)?,
))
}) {
Some((left_load, right_load)) => compare_decode_load(left_load, right_load)
@@ -532,18 +762,54 @@ pub(crate) fn compare_decode_pressure(
}
}
fn compare_decode_load(left: &EngineWorkerLoad, right: &EngineWorkerLoad) -> Ordering {
fn compare_decode_load(left: &NativeCacheWorkerLoad, right: &NativeCacheWorkerLoad) -> Ordering {
let kv_usage = match (left.max_total_num_tokens, right.max_total_num_tokens) {
(left_cap, right_cap) if left_cap > 0 && right_cap > 0 => u128::from(left.num_tokens)
(left_cap, right_cap) if left_cap > 0 && right_cap > 0 => u128::from(left.num_used_tokens)
.saturating_mul(u128::from(right_cap))
.cmp(&u128::from(right.num_tokens).saturating_mul(u128::from(left_cap))),
.cmp(&u128::from(right.num_used_tokens).saturating_mul(u128::from(left_cap))),
_ => Ordering::Equal,
};
left.num_waiting_reqs
.cmp(&right.num_waiting_reqs)
.then_with(|| left.num_running_reqs.cmp(&right.num_running_reqs))
.then(kv_usage)
.then_with(|| left.num_tokens.cmp(&right.num_tokens))
.then_with(|| left.num_used_tokens.cmp(&right.num_used_tokens))
}
fn pressure_guard_prefers_backup(
primary: &Arc<Worker>,
backup: &Arc<Worker>,
hints: &GuardHints,
snapshot: &EngineLoadSnapshot,
) -> bool {
if !hints.enable_pressure_guard {
return false;
}
let (Some(primary_load), Some(backup_load)) = (
snapshot.fresh_native_cache_load_for_url(&primary.url),
snapshot.fresh_native_cache_load_for_url(&backup.url),
) else {
return false;
};
if let Some(absolute_threshold_ms) = hints.pressure_abs_threshold_ms.filter(|_| {
primary_load.estimated_prefill_queue_ms.is_some()
&& backup_load.estimated_prefill_queue_ms.is_some()
}) {
let primary_ms = primary_load
.estimated_prefill_queue_ms
.expect("availability was checked");
let backup_ms = backup_load
.estimated_prefill_queue_ms
.expect("availability was checked");
return primary_ms - backup_ms > absolute_threshold_ms
&& primary_ms > backup_ms * hints.pressure_rel_threshold;
}
primary_load
.num_waiting_uncached_tokens
.saturating_sub(backup_load.num_waiting_uncached_tokens)
> hints.pressure_abs_threshold_tokens
&& primary_load.num_waiting_uncached_tokens as f64
> backup_load.num_waiting_uncached_tokens as f64 * hints.pressure_rel_threshold
}
#[cfg(test)]
@@ -563,18 +829,23 @@ mod tests {
}
fn snapshot(entries: &[(&Arc<Worker>, u64, u64, u64, u64)]) -> EngineLoadSnapshot {
EngineLoadSnapshot::from_workers(
EngineLoadSnapshot::from_native_cache_workers(
7,
entries
.iter()
.map(|(worker, running, waiting, used, capacity)| {
(
worker.url.clone(),
EngineWorkerLoad {
NativeCacheWorkerLoad {
num_running_reqs: *running,
num_waiting_reqs: *waiting,
num_tokens: *used,
num_waiting_uncached_tokens: *waiting,
num_used_tokens: *used,
num_total_tokens: *used,
max_total_num_tokens: *capacity,
max_running_requests: 64,
prefill_throughput_tokens_per_s: None,
estimated_prefill_queue_ms: None,
captured_at: Instant::now(),
},
)
@@ -589,7 +860,7 @@ mod tests {
let unknown = worker("unknown");
let workers = vec![Arc::clone(&full), Arc::clone(&unknown)];
let range = CandidateRange::global(&workers);
let loads = snapshot(&[(&full, 0, 0, 90, 100), (&unknown, 0, 0, 90, 0)]);
let loads = snapshot(&[(&full, 0, 0, 90, 100), (&unknown, 0, 0, 0, 1_000)]);
assert!(resolve_prefill(
&range,
@@ -600,7 +871,7 @@ mod tests {
.is_some());
assert_eq!(
resolve_prefill(&range, &SelectionProposal::primary(full), 20, &loads)
.expect("fallback selects unknown-capacity worker")
.expect("fallback selects the admitted worker")
.selected
.id,
unknown.id
@@ -669,4 +940,51 @@ mod tests {
let _guard = left.load_guard();
assert!(compare_prefill_pressure(&left, &right, None).is_gt());
}
#[test]
fn complete_monitor_pressure_guard_overrides_a_near_cache_gain() {
let congested = worker("congested");
let idle = worker("idle");
let proposal = CacheCandidateProposal {
candidates: vec![
CacheCandidate {
worker: Arc::clone(&congested),
matched_prefix_tokens: 90,
uncached_tokens: 10,
candidate_range_id: "global".into(),
max_pending_prefill_tokens: None,
},
CacheCandidate {
worker: Arc::clone(&idle),
matched_prefix_tokens: 80,
uncached_tokens: 20,
candidate_range_id: "global".into(),
max_pending_prefill_tokens: None,
},
],
cache_switch_margin_tokens: 32,
enable_pressure_guard: true,
pressure_abs_threshold_tokens: 100,
pressure_abs_threshold_ms: None,
pressure_rel_threshold: 1.5,
};
let loads = snapshot(&[
(&congested, 1, 1_000, 10, 10_000),
(&idle, 1, 10, 10, 10_000),
]);
let resolution = resolve_cache_candidates(&proposal, 100, &loads);
assert_eq!(
resolution
.decision
.expect("the idle candidate remains admitted")
.selected
.id,
idle.id
);
assert_eq!(resolution.prefill_pressure_source, "native_queue_tokens");
assert_eq!(resolution.admission_rejected_candidates, 0);
assert_eq!(resolution.pressure_guard_compared_pairs, 1);
assert_eq!(resolution.pressure_guard_overrides, 1);
}
}
@@ -0,0 +1,305 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Builds ordered candidate domains from request shape, SLO profile, and rank.
use crate::config::{BucketConfig, BucketSpec, BucketStage, SloBucketPolicy};
use crate::policies::admission::CandidateDomain;
use crate::policies::CacheCandidate;
use crate::workers::Worker;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
/// Request fields used for bucket selection.
#[derive(Debug, Clone, Copy)]
pub struct BucketRequest {
pub input_tokens: u64,
pub expected_peak_sequence_tokens: Option<u64>,
pub ttft_slo_ms: Option<u64>,
pub tps_slo: Option<f64>,
}
#[derive(Debug, Clone)]
pub struct BucketSelector {
config: Option<BucketConfig>,
/// Precomputed only for buckets above the measured scan/hash crossover.
member_ids: HashMap<String, HashSet<String>>,
}
/// Measured crossover: SipHash costs more than a few short string comparisons.
const MEMBER_SCAN_MAX: usize = 4;
impl BucketSelector {
pub fn new(config: Option<BucketConfig>) -> Self {
let member_ids = config
.as_ref()
.map(|config| {
config
.buckets
.iter()
.filter(|spec| spec.worker_ids.len() > MEMBER_SCAN_MAX)
.map(|spec| {
(
spec.id.clone(),
spec.worker_ids.iter().cloned().collect::<HashSet<_>>(),
)
})
.collect()
})
.unwrap_or_default();
Self { config, member_ids }
}
pub fn is_enabled(&self) -> bool {
self.config.is_some()
}
pub fn prefill_domains(
&self,
workers: &[Arc<Worker>],
request: BucketRequest,
) -> Vec<CandidateDomain> {
let Some(config) = &self.config else {
return vec![CandidateDomain::global_prefill(workers)];
};
self.ordered_specs(
BucketStage::Prefill,
config.ttft_slo_policy,
|spec| prefill_compatible(spec, request.input_tokens),
|spec| ttft_eligible(spec, request.ttft_slo_ms),
)
.into_iter()
.filter_map(|spec| {
let members = self.members(workers, spec);
(!members.is_empty()).then(|| {
CandidateDomain::bucket_prefill(
spec.id.clone(),
members,
spec.max_pending_prefill_tokens,
)
})
})
.collect()
}
pub fn decode_domains(
&self,
workers: &[Arc<Worker>],
request: BucketRequest,
) -> Vec<CandidateDomain> {
let Some(config) = &self.config else {
return vec![CandidateDomain::global_decode(workers)];
};
// Keep a global decode domain when no decode bucket is configured.
if !config
.buckets
.iter()
.any(|spec| spec.stage == BucketStage::Decode)
{
return vec![CandidateDomain::global_decode(workers)];
}
self.ordered_specs(
BucketStage::Decode,
config.tps_slo_policy,
|spec| {
decode_compatible(
spec,
request.input_tokens,
request.expected_peak_sequence_tokens,
)
},
|spec| tps_eligible(spec, request.tps_slo),
)
.into_iter()
.filter_map(|spec| {
let members = self.members(workers, spec);
(!members.is_empty()).then(|| CandidateDomain::bucket_decode(spec.id.clone(), members))
})
.collect()
}
/// Maps global Indexer candidates to prefill buckets using `E` as the workload.
pub fn bind_prefill_cache_candidate(
&self,
mut candidate: CacheCandidate,
request: BucketRequest,
) -> Option<CacheCandidate> {
let Some(config) = &self.config else {
candidate.candidate_range_id = "global".to_string();
candidate.max_pending_prefill_tokens = None;
return Some(candidate);
};
let spec = config.buckets.iter().find(|spec| {
spec.stage == BucketStage::Prefill
&& self.contains(spec, &candidate.worker.id.0)
&& within(
candidate.uncached_tokens,
spec.min_extend_tokens,
spec.max_extend_tokens,
)
&& spec
.max_context_tokens
.is_none_or(|max_context| request.input_tokens <= max_context)
&& (config.ttft_slo_policy != SloBucketPolicy::SloFirst
|| ttft_eligible(spec, request.ttft_slo_ms))
})?;
candidate.candidate_range_id = spec.id.clone();
candidate.max_pending_prefill_tokens = spec.max_pending_prefill_tokens;
Some(candidate)
}
/// Finds the prefill bucket containing a global session primary.
pub fn prefill_affinity_domain(
&self,
workers: &[Arc<Worker>],
primary: &Arc<Worker>,
request: BucketRequest,
) -> Option<CandidateDomain> {
let config = self.config.as_ref()?;
let spec = config.buckets.iter().find(|spec| {
spec.stage == BucketStage::Prefill
&& self.contains(spec, &primary.id.0)
&& spec
.max_context_tokens
.is_none_or(|max_context| request.input_tokens <= max_context)
&& (config.ttft_slo_policy != SloBucketPolicy::SloFirst
|| ttft_eligible(spec, request.ttft_slo_ms))
})?;
let members = self.members(workers, spec);
members
.iter()
.any(|worker| worker.id == primary.id)
.then(|| {
CandidateDomain::bucket_prefill(
spec.id.clone(),
members,
spec.max_pending_prefill_tokens,
)
})
}
fn contains(&self, spec: &BucketSpec, worker_id: &str) -> bool {
if spec.worker_ids.len() <= MEMBER_SCAN_MAX {
return spec.worker_ids.iter().any(|id| id == worker_id);
}
self.member_ids
.get(&spec.id)
.is_some_and(|ids| ids.contains(worker_id))
}
fn members(&self, workers: &[Arc<Worker>], spec: &BucketSpec) -> Vec<Arc<Worker>> {
if spec.worker_ids.len() <= MEMBER_SCAN_MAX {
return workers
.iter()
.filter(|worker| spec.worker_ids.iter().any(|id| id == &worker.id.0))
.cloned()
.collect();
}
let ids = self
.member_ids
.get(&spec.id)
.expect("large bucket member index is built with the config");
workers
.iter()
.filter(|worker| ids.contains(&worker.id.0))
.cloned()
.collect()
}
fn ordered_specs(
&self,
stage: BucketStage,
slo_policy: SloBucketPolicy,
compatible: impl Fn(&BucketSpec) -> bool,
slo_eligible: impl Fn(&BucketSpec) -> bool,
) -> Vec<&BucketSpec> {
let Some(config) = &self.config else {
return Vec::new();
};
let mut compatible_specs: Vec<&BucketSpec> = config
.buckets
.iter()
.filter(|spec| spec.stage == stage && compatible(spec))
.collect();
compatible_specs.sort_by(|left, right| {
left.rank
.cmp(&right.rank)
.then_with(|| left.id.cmp(&right.id))
});
if slo_policy == SloBucketPolicy::Disabled {
return compatible_specs;
}
let mut eligible = Vec::new();
let mut degraded = Vec::new();
for spec in compatible_specs {
if slo_eligible(spec) {
eligible.push(spec);
} else {
degraded.push(spec);
}
}
match slo_policy {
SloBucketPolicy::Disabled => unreachable!("handled before SLO partitioning"),
SloBucketPolicy::SloFirst => {
eligible.extend(degraded);
eligible
}
SloBucketPolicy::BestEffort => {
// Best effort prefers a bucket without an SLO tier.
degraded.extend(eligible);
degraded
}
}
}
}
fn prefill_compatible(spec: &BucketSpec, input_tokens: u64) -> bool {
// With no cache hit, E equals L.
within(input_tokens, spec.min_extend_tokens, spec.max_extend_tokens)
&& spec
.max_context_tokens
.is_none_or(|max_context| input_tokens <= max_context)
}
fn decode_compatible(
spec: &BucketSpec,
input_tokens: u64,
expected_peak_sequence_tokens: Option<u64>,
) -> bool {
let Some(expected_peak_sequence_tokens) = expected_peak_sequence_tokens else {
// Unknown output length can only use a catch-all decode bucket.
return spec.min_sequence_tokens.is_none()
&& spec.max_sequence_tokens.is_none()
&& spec
.max_context_tokens
.is_none_or(|max_context| input_tokens <= max_context);
};
within(
expected_peak_sequence_tokens,
spec.min_sequence_tokens,
spec.max_sequence_tokens,
) && spec
.max_context_tokens
.is_none_or(|max_context| expected_peak_sequence_tokens <= max_context)
}
fn within(value: u64, min: Option<u64>, max: Option<u64>) -> bool {
min.is_none_or(|min| value >= min) && max.is_none_or(|max| value <= max)
}
fn ttft_eligible(spec: &BucketSpec, request_ttft_slo_ms: Option<u64>) -> bool {
let Some(request_ttft_slo_ms) = request_ttft_slo_ms else {
return true;
};
spec.ttft_p95_at_capacity_ms
.is_some_and(|p95| p95 <= request_ttft_slo_ms)
}
fn tps_eligible(spec: &BucketSpec, request_tps_slo: Option<f64>) -> bool {
let Some(request_tps_slo) = request_tps_slo else {
return true;
};
spec.tps_p05_at_capacity
.is_some_and(|p05| p05 >= request_tps_slo)
}
@@ -72,6 +72,13 @@ impl CacheAwarePolicy {
});
}
if let Some((selector, request)) = ctx.prefill_cache_bucket() {
candidates = candidates
.into_iter()
.filter_map(|candidate| selector.bind_prefill_cache_candidate(candidate, request))
.collect();
}
let limit = self.candidate_limit(workers.len());
if limit == 0 {
return None;
@@ -93,6 +100,10 @@ impl CacheAwarePolicy {
Some(CacheCandidateProposal {
candidates,
cache_switch_margin_tokens: self.config.cache_switch_margin_tokens,
enable_pressure_guard: self.config.pressure_guard,
pressure_abs_threshold_tokens: self.config.pressure_abs_threshold_tokens,
pressure_abs_threshold_ms: self.config.pressure_abs_threshold_ms,
pressure_rel_threshold: self.config.pressure_rel_threshold,
})
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,164 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Decode policy extension point, independent of prefill affinity.
use crate::config::DecodePolicyKind;
use crate::policies::admission::{
compare_decode_pressure, resolve_decode, CandidateDomain, DecisionReason, FinalDecision,
RoutingStage,
};
use crate::policies::engine_load::EngineLoadSnapshot;
use crate::policies::registry::select_decode_with_affinity;
use crate::policies::{ProposalKind, SelectionProposal};
use rand::Rng;
use std::sync::Arc;
#[derive(Debug, Default)]
pub struct DecodeSelectionContext<'a> {
load_snapshot: Option<&'a EngineLoadSnapshot>,
prefill_url: Option<&'a str>,
}
impl<'a> DecodeSelectionContext<'a> {
pub fn new() -> Self {
Self {
load_snapshot: None,
prefill_url: None,
}
}
/// Engine load snapshot captured at request ingress.
pub fn with_load_snapshot(mut self, load_snapshot: &'a EngineLoadSnapshot) -> Self {
self.load_snapshot = Some(load_snapshot);
self
}
pub fn load_snapshot(&self) -> Option<&EngineLoadSnapshot> {
self.load_snapshot
}
/// Prefill URL used by `legacy_host_affinity`.
pub fn with_prefill_url(mut self, prefill_url: &'a str) -> Self {
self.prefill_url = Some(prefill_url);
self
}
pub fn prefill_url(&self) -> Option<&str> {
self.prefill_url
}
}
pub trait DecodePolicy: Send + Sync + std::fmt::Debug {
fn propose(
&self,
domain: &CandidateDomain,
ctx: &DecodeSelectionContext<'_>,
) -> Option<SelectionProposal>;
}
/// Resolves decode admission and degrades to Power-of-Two when capacity is exhausted.
pub fn resolve_decode_with_capacity_fallback(
domain: &CandidateDomain,
proposal: &SelectionProposal,
request_kv_tokens: u64,
snapshot: &EngineLoadSnapshot,
) -> Option<FinalDecision> {
if let Some(decision) = resolve_decode(domain, proposal, request_kv_tokens, snapshot) {
return Some(decision);
}
if domain.stage != RoutingStage::Decode
|| !domain
.workers
.iter()
.any(|worker| worker.id == proposal.primary.id)
{
return None;
}
let fallback = DecodePowerOfTwoPolicy::new().propose(
domain,
&DecodeSelectionContext::new().with_load_snapshot(snapshot),
)?;
Some(FinalDecision {
selected: fallback.primary,
primary: Arc::clone(&proposal.primary),
backup: proposal
.backup
.as_ref()
.filter(|backup| domain.workers.iter().any(|worker| worker.id == backup.id))
.cloned(),
reason: DecisionReason::CapacityFallbackPowerOfTwo,
candidate_range_id: domain.id.clone(),
load_snapshot_version: snapshot.version,
})
}
/// Samples two workers from a decode domain and orders them by decode pressure.
#[derive(Debug, Default)]
pub struct DecodePowerOfTwoPolicy;
impl DecodePowerOfTwoPolicy {
pub fn new() -> Self {
Self
}
}
impl DecodePolicy for DecodePowerOfTwoPolicy {
fn propose(
&self,
domain: &CandidateDomain,
ctx: &DecodeSelectionContext<'_>,
) -> Option<SelectionProposal> {
match domain.workers.len() {
0 => None,
1 => Some(
SelectionProposal::primary(Arc::clone(&domain.workers[0]))
.with_kind(ProposalKind::PowerOfTwo),
),
len => {
let mut rng = rand::thread_rng();
let i = rng.gen_range(0..len);
let mut j = rng.gen_range(0..len - 1);
if j >= i {
j += 1;
}
let left = &domain.workers[i];
let right = &domain.workers[j];
let (primary, backup) =
if compare_decode_pressure(left, right, ctx.load_snapshot()).is_gt() {
(Arc::clone(right), Arc::clone(left))
} else {
(Arc::clone(left), Arc::clone(right))
};
Some(
SelectionProposal::with_backup(primary, backup)
.with_kind(ProposalKind::PowerOfTwo),
)
}
}
}
}
/// Compatibility policy for legacy same-host PD decode selection.
#[derive(Debug, Default)]
pub struct LegacyHostAffinityDecodePolicy;
impl DecodePolicy for LegacyHostAffinityDecodePolicy {
fn propose(
&self,
domain: &CandidateDomain,
ctx: &DecodeSelectionContext<'_>,
) -> Option<SelectionProposal> {
let prefill_url = ctx.prefill_url()?;
select_decode_with_affinity(prefill_url, &domain.workers).map(SelectionProposal::primary)
}
}
/// Builds a decode policy scoped to one role.
pub fn build_decode_policy(kind: DecodePolicyKind) -> Box<dyn DecodePolicy> {
match kind {
DecodePolicyKind::PowerOfTwo => Box::new(DecodePowerOfTwoPolicy::new()),
DecodePolicyKind::LegacyHostAffinity => Box::new(LegacyHostAffinityDecodePolicy),
}
}
@@ -6,11 +6,8 @@
//! Workers publish a [`LoadStat`] gauge on their dedicated load socket (see
//! `python/sglang/srt/managers/scheduler_components/load_publisher.py`). The
//! load subscriber routes those into this table, keyed per
//! `(worker_url, dp_rank)`; the
//! cache-aware-zmq policy reads the freshest aggregate per worker as a
//! truthful load signal, falling back to the router-side in-flight counter
//! when no fresh snapshot exists (cold start, stale publisher, or a worker
//! that predates load publishing).
//! `(worker_url, dp_rank)`. Request handling captures the freshest complete
//! aggregate and falls back to Router-local load when it is unavailable.
//!
//! Load is a *gauge*, not a delta: last value wins, no sequence/replay
//! semantics. Entries older than [`EngineLoadTable::freshness`] are ignored.
@@ -25,15 +22,27 @@ use dashmap::{DashMap, DashSet};
use serde::de::{self, Deserializer, IgnoredAny, SeqAccess, Visitor};
use serde::Deserialize;
/// Per-rank load fields consumed by native Cache-Aware.
///
/// Short frames cannot drive admission or pressure guards, so native
/// `cache_aware` falls back to Router-local load.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NativeCacheRankLoad {
pub num_waiting_uncached_tokens: u64,
pub num_total_tokens: u64,
pub max_running_requests: u64,
pub total_prefill_uncached_tokens: u64,
pub total_prefill_busy_us: u64,
}
/// Per-scheduler runtime load snapshot. Mirrors the Python `LoadStat` in
/// `managers/scheduler_components/load_publisher.py`, published on the
/// worker's dedicated load socket (separate from KV-cache events).
///
/// Wire shape (msgspec `tag=True` + `array_like`):
/// `["LoadStat", num_running_reqs, num_waiting_reqs, num_tokens,
/// max_total_num_tokens, attn_dp_rank?]`. We read the four counts and ignore
/// any trailing fields (`attn_dp_rank` — the router keys load by the
/// subscriber's socket rank, not the payload).
/// The stable prefix remains `["LoadStat", running, waiting, used_tokens,
/// max_tokens, attn_dp_rank]`; V4 appends the V3 native Cache-Aware fields.
/// Older publishers therefore decode successfully with `native_cache=None`,
/// which deliberately excludes them from monitor-backed admission/guard.
#[derive(Debug, Clone, PartialEq)]
pub struct LoadStat {
/// Requests currently running on the engine.
@@ -44,9 +53,15 @@ pub struct LoadStat {
pub num_tokens: u64,
/// KV-cache token capacity; 0 when unknown.
pub max_total_num_tokens: u64,
/// V3 native Cache-Aware semantics. `None` means the publisher is an old
/// four-field #34608 producer or sent a truncated extension.
pub native_cache: Option<NativeCacheRankLoad>,
}
/// Aggregated, usable Engine load for one Worker at a fixed instant.
/// Engine load for one worker captured at a fixed point in time.
///
/// The four #34608 fields are summed across DP ranks. `captured_at` retains
/// the oldest rank timestamp so later local dispatches can be added.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EngineWorkerLoad {
pub num_running_reqs: u64,
@@ -56,11 +71,33 @@ pub struct EngineWorkerLoad {
pub captured_at: Instant,
}
/// Immutable Engine-load view captured once at request ingress.
/// Complete ZMQ monitor aggregate used by native Cache-Aware.
///
/// Prefill throughput and queue time require two monotonic samples from every
/// DP rank. Initial samples and counter resets leave both values unavailable.
#[derive(Debug, Clone, PartialEq)]
pub struct NativeCacheWorkerLoad {
pub num_running_reqs: u64,
pub num_waiting_reqs: u64,
pub num_waiting_uncached_tokens: u64,
pub num_used_tokens: u64,
pub num_total_tokens: u64,
pub max_total_num_tokens: u64,
pub max_running_requests: u64,
pub prefill_throughput_tokens_per_s: Option<f64>,
pub estimated_prefill_queue_ms: Option<f64>,
pub captured_at: Instant,
}
/// Immutable engine load view captured once at request ingress.
///
/// Keys are worker URLs used for dispatch. Missing, stale, or rank-incomplete
/// workers are omitted and must use Router-local active load.
#[derive(Debug, Clone, Default)]
pub struct EngineLoadSnapshot {
pub version: u64,
workers: HashMap<String, EngineWorkerLoad>,
native_cache_workers: HashMap<String, NativeCacheWorkerLoad>,
}
impl EngineLoadSnapshot {
@@ -68,9 +105,50 @@ impl EngineLoadSnapshot {
self.workers.get(worker_url)
}
/// Builds a view from already validated Worker data for tests and offline checks.
/// Returns only complete, fresh native Cache-Aware monitor data.
pub fn fresh_native_cache_load_for_url(
&self,
worker_url: &str,
) -> Option<&NativeCacheWorkerLoad> {
self.native_cache_workers.get(worker_url)
}
/// Builds a view from worker data that already passed freshness and rank checks.
/// Production requests should use [`EngineLoadTable::capture_snapshot`].
pub fn from_workers(version: u64, workers: HashMap<String, EngineWorkerLoad>) -> Self {
Self { version, workers }
Self {
version,
workers,
native_cache_workers: HashMap::new(),
}
}
/// Builds a test snapshot from complete native monitor data.
/// Production requests must use [`EngineLoadTable::capture_snapshot`].
pub fn from_native_cache_workers(
version: u64,
workers: HashMap<String, NativeCacheWorkerLoad>,
) -> Self {
let basic = workers
.iter()
.map(|(url, load)| {
(
url.clone(),
EngineWorkerLoad {
num_running_reqs: load.num_running_reqs,
num_waiting_reqs: load.num_waiting_reqs,
num_tokens: load.num_used_tokens,
max_total_num_tokens: load.max_total_num_tokens,
captured_at: load.captured_at,
},
)
})
.collect();
Self {
version,
workers: basic,
native_cache_workers: workers,
}
}
}
@@ -116,12 +194,46 @@ impl<'de> Deserialize<'de> for LoadStat {
let max_total_num_tokens: u64 = seq
.next_element()?
.ok_or_else(|| de::Error::missing_field("max_total_num_tokens"))?;
// `attn_dp_rank` is informational: the subscriber's socket
// rank is authoritative for aggregation. Keep accepting null
// and integer values from both old and new publishers.
let _attn_dp_rank: Option<IgnoredAny> = seq.next_element()?;
// The extension is deliberately all-or-nothing. A four-field
// #34608 message remains valid for lightweight queue routing,
// but a partial semantic tail is not valid monitor data.
let native_cache = match seq.next_element::<u64>()? {
None => None,
Some(num_waiting_uncached_tokens) => {
let num_total_tokens = seq
.next_element()?
.ok_or_else(|| de::Error::missing_field("num_total_tokens"))?;
let max_running_requests = seq
.next_element()?
.ok_or_else(|| de::Error::missing_field("max_running_requests"))?;
let total_prefill_uncached_tokens =
seq.next_element()?.ok_or_else(|| {
de::Error::missing_field("total_prefill_uncached_tokens")
})?;
let total_prefill_busy_us = seq
.next_element()?
.ok_or_else(|| de::Error::missing_field("total_prefill_busy_us"))?;
Some(NativeCacheRankLoad {
num_waiting_uncached_tokens,
num_total_tokens,
max_running_requests,
total_prefill_uncached_tokens,
total_prefill_busy_us,
})
}
};
while seq.next_element::<IgnoredAny>()?.is_some() {}
Ok(LoadStat {
num_running_reqs,
num_waiting_reqs,
num_tokens,
max_total_num_tokens,
native_cache,
})
}
}
@@ -143,12 +255,15 @@ const DEFAULT_FRESHNESS: Duration = Duration::from_secs(5);
#[derive(Debug, Clone)]
struct LoadEntry {
load: LoadStat,
previous_native_cache: Option<NativeCacheRankLoad>,
at: Instant,
}
/// Per-`(worker_url, dp_rank)` engine-reported load. Written by the load
/// subscriber pump, read by the cache-aware-zmq policy. Shared out of
/// [`super::kv_events::index::KvEventIndex`] the same way the hash tree is.
type NativeRankObservation = (LoadStat, Option<NativeCacheRankLoad>, bool, Instant);
type NativeWorkerObservations = HashMap<u32, NativeRankObservation>;
/// Per-`(worker_url, dp_rank)` engine-reported load, written by the load
/// subscriber pump and captured once at request ingress.
#[derive(Debug)]
pub struct EngineLoadTable {
by_rank: DashMap<(String, u32), LoadEntry>,
@@ -182,8 +297,19 @@ impl EngineLoadTable {
/// Record the latest load for one `(worker_url, dp_rank)`.
pub fn set(&self, url: &str, dp_rank: u32, load: LoadStat, at: Instant) {
self.by_rank
.insert((url.to_string(), dp_rank), LoadEntry { load, at });
let key = (url.to_string(), dp_rank);
let previous_native_cache = self
.by_rank
.get(&key)
.and_then(|entry| entry.load.native_cache.clone());
self.by_rank.insert(
key,
LoadEntry {
load,
previous_native_cache,
at,
},
);
self.version.fetch_add(1, Ordering::Relaxed);
}
@@ -194,9 +320,7 @@ impl EngineLoadTable {
}
}
/// Number of workers expected to publish load. Compared against the size
/// of [`Self::snapshot_fresh`] to surface a dead/misconfigured publisher
/// (expected > 0 but no fresh snapshots) in logs.
/// Number of workers expected to publish load.
pub fn expected_count(&self) -> usize {
self.expected
.iter()
@@ -205,33 +329,14 @@ impl EngineLoadTable {
.len()
}
/// Shared accumulation pass behind [`Self::snapshot_fresh`] and
/// [`Self::capture_snapshot`]. It produces the #34608 fields summed across
/// ranks and the OLDEST snapshot timestamp — **but only for workers whose
/// Shared accumulation pass behind [`Self::capture_snapshot`]. It sums
/// fields across ranks and keeps the oldest snapshot timestamp, but only for workers whose
/// every advertised rank is present and fresh**. A missing or stale rank is
/// omitted, so the caller falls back to its own load signal. (Summing
/// only the fresh ranks would make a worker whose other ranks went silent
/// look misleadingly idle and draw *more* traffic.) Callers that never
/// registered expected ranks retain the legacy all-known-ranks rule.
/// `snapshot_fresh` and any other consumer walking this same pass can
/// never disagree with each other about which workers count as fresh.
///
/// The oldest (not newest) rank's timestamp is deliberately what's kept
/// alongside the depth: a caller using it as a "dispatches not yet
/// reflected in this number" cutoff (see
/// `crate::policies::cache_aware_zmq::WorkerLoads::load_of`) needs a
/// bound that never treats an unreported dispatch as already-covered —
/// the freshest rank's timestamp could do exactly that for whichever
/// rank published less recently. This conservatism is one-sided, not
/// free: for a multi-rank worker with skewed publish times, a dispatch
/// that landed on (and was already reported by) the FRESHER rank can
/// get re-added by the caller's cutoff-based correction anyway, since
/// that correction has no way to attribute a dispatch to a specific
/// rank. That's an accepted, bounded over-count (it biases the wrong
/// direction relative to the under-count this method exists to avoid,
/// not a correctness hole) rather than something this method can close
/// on its own — closing it would require per-rank dispatch attribution,
/// which the router-side slot tracking below doesn't have.
/// registered expected ranks retain the all-known-ranks rule. The oldest
/// timestamp represents the freshness of the complete aggregate.
fn fresh_worker_loads(&self, now: Instant) -> HashMap<String, EngineWorkerLoad> {
// url -> rank -> (reported load, fresh, timestamp).
let mut observed: HashMap<String, HashMap<u32, (LoadStat, bool, Instant)>> = HashMap::new();
@@ -293,42 +398,127 @@ impl EngineLoadTable {
.collect()
}
/// Captures the immutable view consumed by all routing decisions in one request.
pub fn capture_snapshot(&self, now: Instant) -> EngineLoadSnapshot {
EngineLoadSnapshot {
version: self.version.load(Ordering::Acquire),
workers: self.fresh_worker_loads(now),
}
}
pub(crate) fn fresh_worker_state(&self, now: Instant) -> HashMap<String, (usize, Instant)> {
self.fresh_worker_loads(now)
.into_iter()
.map(|(url, load)| {
/// Aggregates complete native Cache-Aware monitor data.
///
/// Every rank must be fresh, capacity-valid, and include the #34608
/// extension. Otherwise the worker is omitted from monitor-backed guards.
fn fresh_native_cache_worker_loads(
&self,
now: Instant,
) -> HashMap<String, NativeCacheWorkerLoad> {
let mut observed: HashMap<String, NativeWorkerObservations> = HashMap::new();
for entry in self.by_rank.iter() {
let at = entry.value().at;
let fresh = now.saturating_duration_since(at) <= self.freshness;
observed.entry(entry.key().0.clone()).or_default().insert(
entry.key().1,
(
url,
entry.value().load.clone(),
entry.value().previous_native_cache.clone(),
fresh,
at,
),
);
}
let mut expected: HashMap<String, HashSet<u32>> = HashMap::new();
for entry in self.expected.iter() {
expected
.entry(entry.key().0.clone())
.or_default()
.insert(entry.key().1);
}
let workers: HashSet<String> = observed.keys().chain(expected.keys()).cloned().collect();
workers
.into_iter()
.filter_map(|url| {
let ranks = observed.get(&url)?;
let required: Vec<u32> = match expected.get(&url) {
Some(expected_ranks) => expected_ranks.iter().copied().collect(),
None => ranks.keys().copied().collect(),
};
let mut num_running_reqs = 0u64;
let mut num_waiting_reqs = 0u64;
let mut num_waiting_uncached_tokens = 0u64;
let mut num_used_tokens = 0u64;
let mut num_total_tokens = 0u64;
let mut max_total_num_tokens = 0u64;
let mut max_running_requests = 0u64;
let mut oldest_at = None;
let mut prefill_throughput_tokens_per_s = 0.0f64;
let mut complete_prefill_sample = !required.is_empty();
for rank in required {
let (load, previous, fresh, at) = ranks.get(&rank)?;
let native = load.native_cache.as_ref()?;
if !fresh || load.max_total_num_tokens == 0 || native.max_running_requests == 0
{
return None;
}
num_running_reqs = num_running_reqs.saturating_add(load.num_running_reqs);
num_waiting_reqs = num_waiting_reqs.saturating_add(load.num_waiting_reqs);
num_waiting_uncached_tokens = num_waiting_uncached_tokens
.saturating_add(native.num_waiting_uncached_tokens);
num_used_tokens = num_used_tokens.saturating_add(load.num_tokens);
num_total_tokens = num_total_tokens.saturating_add(native.num_total_tokens);
max_total_num_tokens =
max_total_num_tokens.saturating_add(load.max_total_num_tokens);
max_running_requests =
max_running_requests.saturating_add(native.max_running_requests);
oldest_at = Some(oldest_at.map_or(*at, |oldest: Instant| oldest.min(*at)));
match previous {
Some(previous)
if native.total_prefill_uncached_tokens
> previous.total_prefill_uncached_tokens
&& native.total_prefill_busy_us
> previous.total_prefill_busy_us =>
{
let tokens = native.total_prefill_uncached_tokens
- previous.total_prefill_uncached_tokens;
let busy_us =
native.total_prefill_busy_us - previous.total_prefill_busy_us;
let rate = 1_000_000.0 * tokens as f64 / busy_us as f64;
if rate.is_finite() && rate > 0.0 {
prefill_throughput_tokens_per_s += rate;
} else {
complete_prefill_sample = false;
}
}
_ => complete_prefill_sample = false,
}
}
let prefill_throughput_tokens_per_s =
complete_prefill_sample.then_some(prefill_throughput_tokens_per_s);
let estimated_prefill_queue_ms = prefill_throughput_tokens_per_s
.map(|rate| 1_000.0 * num_waiting_uncached_tokens as f64 / rate);
oldest_at.map(|captured_at| {
(
load.num_running_reqs
.saturating_add(load.num_waiting_reqs)
.try_into()
.unwrap_or(usize::MAX),
load.captured_at,
),
)
url,
NativeCacheWorkerLoad {
num_running_reqs,
num_waiting_reqs,
num_waiting_uncached_tokens,
num_used_tokens,
num_total_tokens,
max_total_num_tokens,
max_running_requests,
prefill_throughput_tokens_per_s,
estimated_prefill_queue_ms,
captured_at,
},
)
})
})
.collect()
}
/// Per worker URL, the summed queue depth (`num_running_reqs +
/// num_waiting_reqs`) across that worker's ranks, for workers whose
/// every advertised rank is fresh. Computed once per selection so per-worker
/// lookups are O(1). See [`Self::fresh_worker_state`] for the freshness
/// gate behind this.
pub fn snapshot_fresh(&self, now: Instant) -> HashMap<String, usize> {
self.fresh_worker_state(now)
.into_iter()
.map(|(url, (depth, _))| (url, depth))
.collect()
/// Captures one immutable view for all routing decisions in a request.
pub fn capture_snapshot(&self, now: Instant) -> EngineLoadSnapshot {
EngineLoadSnapshot {
version: self.version.load(Ordering::Acquire),
workers: self.fresh_worker_loads(now),
native_cache_workers: self.fresh_native_cache_worker_loads(now),
}
}
/// Drop every rank entry (and the expected mark) for a worker. Called on
@@ -355,6 +545,7 @@ mod tests {
num_waiting_reqs: waiting,
num_tokens: 0,
max_total_num_tokens: 0,
native_cache: None,
}
}
@@ -377,15 +568,54 @@ mod tests {
assert!(decode_load_stat(&missing_count).is_err());
}
#[test]
fn load_wire_preserves_the_v3_native_cache_extension_and_accepts_old_short_frames() {
let mut full = Vec::new();
rmp::encode::write_array_len(&mut full, 11).unwrap();
rmp::encode::write_str(&mut full, "LoadStat").unwrap();
for value in [2, 3, 4, 100] {
rmp::encode::write_u64(&mut full, value).unwrap();
}
rmp::encode::write_nil(&mut full).unwrap();
for value in [500, 600, 32, 1_000, 2_000] {
rmp::encode::write_u64(&mut full, value).unwrap();
}
let decoded = decode_load_stat(&full).expect("complete extended LoadStat decodes");
assert_eq!(decoded.num_running_reqs, 2);
assert_eq!(
decoded
.native_cache
.expect("extension must be retained")
.num_waiting_uncached_tokens,
500
);
let mut old = Vec::new();
rmp::encode::write_array_len(&mut old, 6).unwrap();
rmp::encode::write_str(&mut old, "LoadStat").unwrap();
for value in [2, 3, 4, 100] {
rmp::encode::write_u64(&mut old, value).unwrap();
}
rmp::encode::write_nil(&mut old).unwrap();
assert!(
decode_load_stat(&old)
.expect("old #34608 four-field frame remains decodable")
.native_cache
.is_none(),
"short frames must never be promoted to complete native monitor data"
);
}
#[test]
fn sums_queue_depth_across_ranks() {
let t = EngineLoadTable::new();
let now = Instant::now();
t.set("http://w:30000", 0, load(5, 1), now);
t.set("http://w:30000", 1, load(3, 2), now);
let fresh = t.snapshot_fresh(now);
let fresh = t.capture_snapshot(now);
// (5+1) + (3+2) = 11
assert_eq!(fresh.get("http://w:30000").copied(), Some(11));
let load = fresh.fresh_load_for_url("http://w:30000").unwrap();
assert_eq!(load.num_running_reqs + load.num_waiting_reqs, 11);
}
#[test]
@@ -395,7 +625,10 @@ mod tests {
t.set("http://w:30000", 0, load(9, 9), old);
// A read far in the future sees the entry as stale -> worker absent.
let later = old + Duration::from_secs(60);
assert!(!t.snapshot_fresh(later).contains_key("http://w:30000"));
assert!(t
.capture_snapshot(later)
.fresh_load_for_url("http://w:30000")
.is_none());
}
#[test]
@@ -407,8 +640,9 @@ mod tests {
t.set("http://other:30000", 0, load(1, 0), now);
t.forget_worker("http://w:30000");
assert_eq!(t.entry_count(), 1);
assert!(!t.snapshot_fresh(now).contains_key("http://w:30000"));
assert!(t.snapshot_fresh(now).contains_key("http://other:30000"));
let snapshot = t.capture_snapshot(now);
assert!(snapshot.fresh_load_for_url("http://w:30000").is_none());
assert!(snapshot.fresh_load_for_url("http://other:30000").is_some());
}
/// A worker with any stale rank is omitted entirely (not summed over only
@@ -422,7 +656,9 @@ mod tests {
t.set("http://w:30000", 0, load(5, 1), now); // fresh
t.set("http://w:30000", 1, load(9, 9), stale); // stale
assert!(
!t.snapshot_fresh(now).contains_key("http://w:30000"),
t.capture_snapshot(now)
.fresh_load_for_url("http://w:30000")
.is_none(),
"any stale rank must drop the whole worker from the snapshot"
);
}
@@ -435,44 +671,30 @@ mod tests {
t.mark_expected_rank("http://w:30000", 1);
t.set("http://w:30000", 0, load(5, 1), now);
assert!(
!t.snapshot_fresh(now).contains_key("http://w:30000"),
t.capture_snapshot(now)
.fresh_load_for_url("http://w:30000")
.is_none(),
"an advertised rank without a reading must not produce a partial aggregate"
);
t.set("http://w:30000", 1, load(3, 2), now);
assert_eq!(t.snapshot_fresh(now).get("http://w:30000"), Some(&11));
let snapshot = t.capture_snapshot(now);
let load = snapshot.fresh_load_for_url("http://w:30000").unwrap();
assert_eq!(load.num_running_reqs + load.num_waiting_reqs, 11);
}
#[test]
fn fresh_worker_state_picks_the_earliest_rank_timestamp() {
fn capture_snapshot_uses_the_earliest_rank_timestamp() {
let t = EngineLoadTable::new();
let earlier = Instant::now() - Duration::from_secs(2);
let later = earlier + Duration::from_secs(1);
t.set("http://w:30000", 0, load(5, 1), later);
t.set("http://w:30000", 1, load(3, 2), earlier);
let now = later + Duration::from_millis(1);
assert_eq!(
t.fresh_worker_state(now).get("http://w:30000").copied(),
Some((11, earlier)),
"must expose the OLDEST rank's timestamp, not the newest"
);
}
#[test]
fn fresh_worker_state_agrees_with_snapshot_fresh_on_which_workers_are_present() {
let t = EngineLoadTable::with_freshness(Duration::from_secs(5));
let now = Instant::now();
let stale = now - Duration::from_secs(3600);
t.set("http://fresh:30000", 0, load(1, 0), now);
t.set("http://mixed:30000", 0, load(1, 0), now);
t.set("http://mixed:30000", 1, load(1, 0), stale);
let depths = t.snapshot_fresh(now);
let state = t.fresh_worker_state(now);
assert!(depths.contains_key("http://fresh:30000"));
assert!(state.contains_key("http://fresh:30000"));
assert!(!depths.contains_key("http://mixed:30000"));
assert!(!state.contains_key("http://mixed:30000"));
let snapshot = t.capture_snapshot(now);
let load = snapshot.fresh_load_for_url("http://w:30000").unwrap();
assert_eq!(load.num_running_reqs + load.num_waiting_reqs, 11);
assert_eq!(load.captured_at, earlier);
}
#[test]
@@ -486,4 +708,48 @@ mod tests {
t.forget_worker("http://w:30000");
assert_eq!(t.expected_count(), 1);
}
#[test]
fn complete_v3_semantic_samples_derive_prefill_queue_time() {
let t = EngineLoadTable::new();
let first = Instant::now();
let second = first + Duration::from_secs(2);
let mut old = load(2, 3);
old.num_tokens = 16_000;
old.max_total_num_tokens = 32_000;
old.native_cache = Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 1_000,
num_total_tokens: 20_000,
max_running_requests: 64,
total_prefill_uncached_tokens: 10_000,
total_prefill_busy_us: 2_000_000,
});
let mut new = old.clone();
new.num_tokens = 20_000;
let native = new
.native_cache
.as_mut()
.expect("test sample has native-cache extension");
native.num_waiting_uncached_tokens = 4_000;
native.num_total_tokens = 24_000;
native.total_prefill_uncached_tokens = 22_000;
native.total_prefill_busy_us = 4_000_000;
t.mark_expected_rank("http://w:30000", 0);
t.set("http://w:30000", 0, old, first);
t.set("http://w:30000", 0, new, second);
let snapshot = t.capture_snapshot(second);
let worker = snapshot
.fresh_native_cache_load_for_url("http://w:30000")
.expect("complete fresh rank must be usable");
assert_eq!(worker.num_waiting_uncached_tokens, 4_000);
assert_eq!(worker.num_total_tokens, 24_000);
assert_eq!(worker.max_running_requests, 64);
assert_eq!(worker.prefill_throughput_tokens_per_s, Some(6_000.0));
assert_eq!(
worker.estimated_prefill_queue_ms,
Some(666.666_666_666_666_6)
);
}
}
+22 -111
View File
@@ -7,8 +7,6 @@ use crate::config::{
use crate::discovery::ModelId;
use crate::policies::{
cache_aware::CacheAwarePolicy,
cache_aware_zmq::CacheAwareZmqPolicy,
engine_load::EngineLoadTable,
kv_events::{BlockSizeOracle, HashTree},
load_based::LoadBasedPolicy,
power_of_two::PowerOfTwoChoicesPolicy,
@@ -22,7 +20,6 @@ use crate::policies::{
sticky::StickyPolicy,
Policy, PolicyRegistry,
};
use crate::tokenizer::TokenizerRegistry;
use anyhow::{anyhow, Result};
use std::sync::Arc;
use std::time::Duration;
@@ -53,19 +50,10 @@ fn build_sticky(model: &ModelConfig) -> Arc<dyn Policy> {
pub fn build_policy(
model: &ModelConfig,
tree: Arc<HashTree>,
tokenizers: Arc<TokenizerRegistry>,
block_size_oracle: Arc<BlockSizeOracle>,
engine_load: Arc<EngineLoadTable>,
) -> Result<Arc<dyn Policy>> {
validate_eligibility(model)?;
let inner = build_kind(
model.policy,
model,
&tree,
&tokenizers,
&block_size_oracle,
&engine_load,
)?;
let inner = build_kind(model.policy, model, &tree, &block_size_oracle)?;
let Some(elig) = model.eligibility.as_ref().filter(|e| !e.filters.is_empty()) else {
return Ok(inner);
};
@@ -101,30 +89,14 @@ fn build_kind(
kind: PolicyKind,
model: &ModelConfig,
tree: &Arc<HashTree>,
tokenizers: &Arc<TokenizerRegistry>,
block_size_oracle: &Arc<BlockSizeOracle>,
engine_load: &Arc<EngineLoadTable>,
) -> Result<Arc<dyn Policy>> {
let (tree, tokenizers, block_size_oracle) = (
Arc::clone(tree),
Arc::clone(tokenizers),
Arc::clone(block_size_oracle),
);
let (tree, block_size_oracle) = (Arc::clone(tree), Arc::clone(block_size_oracle));
Ok(match kind {
PolicyKind::RoundRobin => Arc::new(RoundRobinPolicy::new()),
PolicyKind::Random => Arc::new(RandomPolicy::new()),
PolicyKind::PowerOfTwo => Arc::new(PowerOfTwoChoicesPolicy::new()),
PolicyKind::LoadBased => Arc::new(LoadBasedPolicy::new()),
PolicyKind::CacheAwareZmq => {
let cache_cfg = model.cache_aware.clone().unwrap_or_default();
Arc::new(CacheAwareZmqPolicy::new(
cache_cfg,
tree,
tokenizers,
block_size_oracle,
Arc::clone(engine_load),
))
}
PolicyKind::SessionAware => Arc::new(SessionAwarePolicy::new(
model.affinity.clone().unwrap_or_default(),
)),
@@ -222,13 +194,6 @@ pub fn build_policy_kind_only(kind: PolicyKind) -> Result<Arc<dyn Policy>> {
PolicyKind::Random => Arc::new(RandomPolicy::new()),
PolicyKind::PowerOfTwo => Arc::new(PowerOfTwoChoicesPolicy::new()),
PolicyKind::LoadBased => Arc::new(LoadBasedPolicy::new()),
PolicyKind::CacheAwareZmq => Arc::new(CacheAwareZmqPolicy::new(
crate::config::CacheAwareConfig::default(),
Arc::new(HashTree::new()),
Arc::new(TokenizerRegistry::default()),
BlockSizeOracle::new(),
EngineLoadTable::new(),
)),
PolicyKind::SessionAware => Arc::new(SessionAwarePolicy::new(
crate::config::AffinityConfig::default(),
)),
@@ -252,34 +217,20 @@ pub fn build_policy_kind_only(kind: PolicyKind) -> Result<Arc<dyn Policy>> {
pub fn build_registry(
cfg: &Config,
tree: Arc<HashTree>,
tokenizers: Arc<TokenizerRegistry>,
block_size_oracle: Arc<BlockSizeOracle>,
engine_load: Arc<EngineLoadTable>,
) -> Result<PolicyRegistry> {
let reg = PolicyRegistry::default();
let m = &cfg.model;
reg.insert(
ModelId(m.id.clone()),
build_policy(
m,
Arc::clone(&tree),
Arc::clone(&tokenizers),
Arc::clone(&block_size_oracle),
Arc::clone(&engine_load),
)?,
build_policy(m, Arc::clone(&tree), Arc::clone(&block_size_oracle))?,
);
Ok(reg)
}
/// Builds a registry with empty cache-aware dependencies.
pub fn build_registry_with_defaults(cfg: &Config) -> Result<PolicyRegistry> {
build_registry(
cfg,
Arc::new(HashTree::new()),
Arc::new(TokenizerRegistry::default()),
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
build_registry(cfg, Arc::new(HashTree::new()), BlockSizeOracle::new())
}
#[cfg(test)]
@@ -386,6 +337,8 @@ mod tests {
id: id.into(),
tokenizer_path: "/tmp/x".into(),
policy,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -403,15 +356,14 @@ mod tests {
#[test]
fn build_policy_kind_only_covers_all_variants() {
for (kind, needs_load_snapshot) in [
(PolicyKind::RoundRobin, false),
(PolicyKind::Random, false),
(PolicyKind::PowerOfTwo, true),
(PolicyKind::LoadBased, true),
(PolicyKind::CacheAwareZmq, true),
(PolicyKind::SessionAware, true),
(PolicyKind::CacheAware, true),
(PolicyKind::Sticky, false),
for (kind, needs_load_snapshot, needs_dispatch_timestamps) in [
(PolicyKind::RoundRobin, false, false),
(PolicyKind::Random, false, false),
(PolicyKind::PowerOfTwo, true, false),
(PolicyKind::LoadBased, true, true),
(PolicyKind::SessionAware, true, false),
(PolicyKind::CacheAware, true, false),
(PolicyKind::Sticky, false, false),
] {
let policy = build_policy_kind_only(kind).unwrap();
assert_eq!(
@@ -419,6 +371,11 @@ mod tests {
needs_load_snapshot,
"{kind:?}"
);
assert_eq!(
policy.needs_dispatch_timestamps(),
needs_dispatch_timestamps,
"{kind:?}"
);
}
assert!(build_policy_kind_only(PolicyKind::FusedScore).is_err());
assert!(build_policy_kind_only(PolicyKind::ScorePolicy).is_err());
@@ -536,54 +493,16 @@ mod tests {
fn registry_assigns_configured_model() {
let cfg = cfg_with_model("qwen", PolicyKind::RoundRobin);
let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(
&cfg,
tree,
tokenizers,
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
.unwrap();
let reg = build_registry(&cfg, tree, BlockSizeOracle::new()).unwrap();
assert!(reg.get(&ModelId("qwen".into())).is_some());
assert!(reg.get(&ModelId("missing".into())).is_none());
}
#[test]
fn cache_aware_zmq_builds_via_factory() {
let cfg = cfg_with_model("modelA", PolicyKind::CacheAwareZmq);
let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(
&cfg,
tree,
tokenizers,
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
.unwrap();
let p = reg.get(&ModelId("modelA".into())).unwrap();
let dbg = format!("{p:?}");
assert!(
dbg.contains("CacheAwareZmqPolicy"),
"expected CacheAwareZmqPolicy debug repr, got: {dbg}",
);
assert!(p.needs_load_snapshot());
}
#[test]
fn load_based_builds_via_factory() {
let cfg = cfg_with_model("modelA", PolicyKind::LoadBased);
let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(
&cfg,
tree,
tokenizers,
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
.unwrap();
let reg = build_registry(&cfg, tree, BlockSizeOracle::new()).unwrap();
let p = reg.get(&ModelId("modelA".into())).unwrap();
let dbg = format!("{p:?}");
assert!(
@@ -597,15 +516,7 @@ mod tests {
fn sticky_builds_via_factory() {
let cfg = cfg_with_model("modelA", PolicyKind::Sticky);
let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(
&cfg,
tree,
tokenizers,
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
.unwrap();
let reg = build_registry(&cfg, tree, BlockSizeOracle::new()).unwrap();
let p = reg.get(&ModelId("modelA".into())).unwrap();
let dbg = format!("{p:?}");
assert!(
@@ -1,8 +1,7 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Process-shared per-(cache-aware-zmq) `block_size`, sourced from the
//! workers themselves.
//! Process-shared KV-cache block size, sourced from the workers.
//!
//! # Why an oracle instead of a config field?
//!
@@ -10,10 +9,6 @@
//! worker uses to publish KV-cache events; otherwise every cache-aware
//! lookup misses silently. The worker advertises its `page_size` via
//! `/server_info` (parsed into [`crate::policies::kv_events::EventConfig::block_size`]).
//! Earlier versions of sgl-router carried a static `block_size` field on
//! `CacheAwareConfig`; nothing reconciled it with the worker-reported
//! value, so a mismatch silently destroyed cache-hit routing.
//!
//! Dynamo's design treats `kv_cache_block_size` as a property of the
//! `ModelDeploymentCard` populated by the worker registrar (see
//! `~/dynamo/components/src/dynamo/sglang/register.py`); a mismatch
@@ -73,9 +68,7 @@ impl BlockSizeOracle {
}
/// Returns the established block size, or `None` if no worker has
/// reported one yet. Routing-time consumers (`CacheAwareZmqPolicy`)
/// fall back to min-load when this is `None`, because they cannot
/// hash a prompt without a block size.
/// reported one yet.
pub fn get(&self) -> Option<u32> {
let v = self.value.load(Ordering::Relaxed);
if v == 0 {
@@ -37,8 +37,8 @@
//! EAGLE-family workers (`is_bigram = is_eagle`) hash KV blocks over
//! overlapping `(t_i, t_{i+1})` token pairs. That path is implemented as a
//! separate [`compute_block_hashes_bigram`] (below) rather than branching
//! inside the non-bigram fast path; `CacheAwareZmqPolicy::select` chooses
//! between the two from the worker-reported bigram flag.
//! inside the non-bigram fast path. Prefix providers choose between them
//! using the worker-reported bigram flag.
use sha2::{Digest, Sha256};
@@ -85,7 +85,7 @@ pub struct KvEventIndex {
/// subscribers for the same worker don't collide.
load_subscribers: Arc<KvEventSubscriberRegistry>,
/// Engine-reported per-worker load, written by the pump from
/// `WorkerEvent::Load` and read by the cache-aware-zmq policy.
/// `WorkerEvent::Load` and captured at request ingress.
engine_load: Arc<EngineLoadTable>,
pump: Mutex<Option<JoinHandle<()>>>,
pump_cancel: CancellationToken,
@@ -103,11 +103,11 @@ pub struct KvEventIndex {
/// may legitimately have a fresh publisher whose sequence numbers
/// restart from 1.
cursors: Arc<Mutex<HashMap<KvWorkerId, i64>>>,
/// Worker-sourced `page_size` shared with the cache-aware-zmq policy.
/// Worker-sourced `page_size` shared with prefix providers.
/// `add_worker` calls `try_set(cfg.block_size)` so the first worker
/// establishes the value; subsequent workers that disagree are
/// rejected (logged + not subscribed). The policy reads it at routing
/// time to size its `compute_block_hashes` call.
/// rejected (logged + not subscribed). Prefix providers read it at routing
/// time to size their `compute_block_hashes` calls.
block_size_oracle: Arc<BlockSizeOracle>,
}
@@ -129,8 +129,8 @@ impl KvEventIndex {
/// Constructor that lets the caller supply a pre-shared
/// [`BlockSizeOracle`]. Production wires this from `AppContext` so
/// the same oracle the index seeds is the one the cache-aware-zmq
/// policy reads at routing time. Tests use this to pre-populate the
/// the same oracle the index seeds is available to prefix providers.
/// Tests use this to pre-populate the
/// oracle and exercise the mismatch-rejection path.
pub fn new_with_http_and_oracle(
http: reqwest::Client,
@@ -186,10 +186,7 @@ impl KvEventIndex {
})
}
/// Shared accessor for the per-process block-size oracle. The
/// `CacheAwareZmqPolicy` (via [`crate::policies::factory`]) holds the
/// same `Arc` so the value the index seeds is the value the policy
/// hashes against.
/// Shared accessor for the per-process block-size oracle.
pub fn block_size_oracle(&self) -> Arc<BlockSizeOracle> {
Arc::clone(&self.block_size_oracle)
}
@@ -201,9 +198,7 @@ impl KvEventIndex {
self.tree.clone()
}
/// Shared accessor for the engine-load table. The `CacheAwareZmqPolicy`
/// (via [`crate::policies::factory`]) holds the same `Arc` and only reads
/// it at selection time. Load *values* are written solely by the pump
/// Shared accessor for the engine-load table. Load values are written solely by the pump
/// (from `LoadStat` events); `add_worker` / `remove_worker` here manage
/// the expected set and per-worker eviction.
pub fn engine_load(&self) -> Arc<EngineLoadTable> {
@@ -616,6 +611,7 @@ mod tests {
num_waiting_reqs: 4,
num_tokens: 0,
max_total_num_tokens: 0,
native_cache: None,
},
})
.await
@@ -623,9 +619,10 @@ mod tests {
drop(tx);
pump.await.unwrap();
let fresh = engine_load.snapshot_fresh(Instant::now());
assert_eq!(fresh.get("http://w1").copied(), Some(12)); // 8 + 4
// Load events must not pollute the cache tree.
let snapshot = engine_load.capture_snapshot(Instant::now());
let load = snapshot.fresh_load_for_url("http://w1").unwrap();
assert_eq!(load.num_running_reqs + load.num_waiting_reqs, 12);
// Load events must not pollute the cache tree.
assert_eq!(tree.node_count(), 0);
}
@@ -911,17 +908,23 @@ mod tests {
num_waiting_reqs: 1,
num_tokens: 0,
max_total_num_tokens: 0,
native_cache: None,
},
now,
);
assert!(index.engine_load().snapshot_fresh(now).contains_key(url));
assert!(index
.engine_load()
.capture_snapshot(now)
.fresh_load_for_url(url)
.is_some());
index.remove_worker(url).await;
assert!(
!index
index
.engine_load()
.snapshot_fresh(Instant::now())
.contains_key(url),
.capture_snapshot(Instant::now())
.fresh_load_for_url(url)
.is_none(),
"remove_worker must clear engine load"
);
assert_eq!(index.engine_load().expected_count(), 0);
@@ -760,11 +760,9 @@ mod tests {
msg
}
/// Wait briefly for the SubSocket to finish its handshake/subscribe.
/// 50ms is empirically enough on localhost without making tests
/// flaky.
/// Allows the local PUB/SUB handshake to complete in concurrent tests.
pub async fn settle() {
tokio::time::sleep(Duration::from_millis(50)).await;
tokio::time::sleep(Duration::from_millis(250)).await;
}
/// Destructure a `WorkerEvent::Batch`, panicking on any other
@@ -1017,7 +1015,7 @@ mod tests {
let mut attempt = 0;
let (pub0, pub1, pub2, base_port) = loop {
attempt += 1;
assert!(attempt < 32, "could not find 3 contiguous free ports");
assert!(attempt < 256, "could not find 3 contiguous free ports");
// Bind PUB at OS-assigned port to learn what's free, then try
// to bind the next two ports explicitly.
@@ -1145,25 +1143,29 @@ mod tests {
.await;
helpers::settle().await;
for (rank, pubsock) in publishers.iter_mut().enumerate() {
pubsock
.send(helpers::build_multipart(
1000 + rank as i64,
helpers::encode_all_blocks_cleared_batch(rank as f64, Some(rank as u32)),
))
.await
.unwrap();
}
let mut by_rank: HashMap<u32, i64> = HashMap::new();
for _ in 0..N {
let event = timeout(Duration::from_millis(500), rx.recv())
.await
.expect("timed out")
.expect("channel closed");
let (worker, seq, _batch) = helpers::expect_batch(event);
assert_eq!(worker.url, worker_url);
by_rank.insert(worker.dp_rank, seq);
for _ in 0..20 {
for (rank, pubsock) in publishers.iter_mut().enumerate() {
pubsock
.send(helpers::build_multipart(
1000 + rank as i64,
helpers::encode_all_blocks_cleared_batch(rank as f64, Some(rank as u32)),
))
.await
.unwrap();
}
for _ in 0..N {
let Ok(Some(event)) = timeout(Duration::from_millis(50), rx.recv()).await else {
break;
};
let (worker, seq, _batch) = helpers::expect_batch(event);
assert_eq!(worker.url, worker_url);
by_rank.insert(worker.dp_rank, seq);
}
if by_rank.len() == N {
break;
}
}
assert_eq!(by_rank.len(), N, "every rank must produce an event");
for rank in 0..N as u32 {
@@ -22,6 +22,10 @@ impl ScoringPolicy for LoadBasedPolicy {
true
}
fn needs_dispatch_timestamps(&self) -> bool {
true
}
/// `1.0` for the least loaded down to `0.0` for the most, min-max scaled to
/// the CURRENT fleet -- relative, not absolute, so it cannot saturate:
/// `1 - load/256` reads a busy fleet as all-`0.0`, tied inside
+118 -27
View File
@@ -3,13 +3,15 @@
pub mod active_load;
pub mod admission;
pub mod buckets;
pub mod cache_aware;
pub mod cache_aware_zmq;
pub mod decode;
pub mod engine_load;
pub mod factory;
pub mod kv_events;
pub mod load_based;
pub mod power_of_two;
pub mod prefix_provider;
pub mod random;
pub mod registry;
pub mod round_robin;
@@ -18,6 +20,7 @@ pub mod session_aware;
pub mod sticky;
use crate::discovery::ModelId;
use crate::policies::buckets::{BucketRequest, BucketSelector};
use crate::policies::engine_load::EngineLoadSnapshot;
use crate::policies::scoring::{EligibilityFilter, ScoringPolicy};
use crate::server::metrics::MetricsRegistry;
@@ -154,6 +157,7 @@ pub struct SelectionContext<'a> {
request_tokens: Option<&'a [u32]>,
external_prefix: Option<&'a ExternalPrefixSignal>,
load_snapshot: Option<&'a EngineLoadSnapshot>,
prefill_cache_bucket: Option<(&'a BucketSelector, BucketRequest)>,
affinity_lookup_enabled: bool,
affinity_assignment_enabled: bool,
}
@@ -170,6 +174,7 @@ impl<'a> SelectionContext<'a> {
request_tokens: None,
external_prefix: None,
load_snapshot: None,
prefill_cache_bucket: None,
affinity_lookup_enabled: true,
affinity_assignment_enabled: true,
}
@@ -190,6 +195,7 @@ impl<'a> SelectionContext<'a> {
request_tokens: None,
external_prefix: None,
load_snapshot: None,
prefill_cache_bucket: None,
affinity_lookup_enabled: true,
affinity_assignment_enabled: true,
}
@@ -201,19 +207,19 @@ impl<'a> SelectionContext<'a> {
self
}
/// Attaches the Session-Aware session ID.
/// Attaches a Session-Aware session ID.
pub fn with_session_id(mut self, session_id: Option<&'a str>) -> Self {
self.session_id = session_id;
self
}
/// Attaches this policy evaluation's candidate range ID.
/// Identifies the candidate domain for this policy call.
pub fn with_candidate_range_id(mut self, candidate_range_id: &'a str) -> Self {
self.candidate_range_id = candidate_range_id;
self
}
/// Attaches the request input-token count.
/// Attaches the request input token count.
pub fn with_input_tokens(mut self, input_tokens: u64) -> Self {
self.input_tokens = Some(input_tokens);
self
@@ -227,12 +233,23 @@ impl<'a> SelectionContext<'a> {
self
}
/// Attaches the Engine Load snapshot captured at request start.
/// Attaches the engine load snapshot captured at request ingress.
pub fn with_load_snapshot(mut self, load_snapshot: &'a EngineLoadSnapshot) -> Self {
self.load_snapshot = Some(load_snapshot);
self
}
/// Cache-Aware uses this binding before Top-K truncation so an
/// incompatible cache holder cannot displace a lower-ranked usable one.
pub fn with_prefill_cache_bucket(
mut self,
selector: &'a BucketSelector,
request: BucketRequest,
) -> Self {
self.prefill_cache_bucket = Some((selector, request));
self
}
/// Disables affinity lookup and assignment.
pub fn without_affinity_lookup(mut self) -> Self {
self.affinity_lookup_enabled = false;
@@ -240,7 +257,7 @@ impl<'a> SelectionContext<'a> {
self
}
/// Keeps affinity lookup but disables assignment writes.
/// Enables affinity lookup without recording new assignments.
pub fn without_affinity_assignment(mut self) -> Self {
self.affinity_assignment_enabled = false;
self
@@ -282,6 +299,10 @@ impl<'a> SelectionContext<'a> {
self.load_snapshot
}
pub fn prefill_cache_bucket(&self) -> Option<(&BucketSelector, BucketRequest)> {
self.prefill_cache_bucket
}
pub fn affinity_lookup_enabled(&self) -> bool {
self.affinity_lookup_enabled
}
@@ -291,36 +312,43 @@ impl<'a> SelectionContext<'a> {
}
}
/// A policy's primary/backup proposal.
/// Primary and backup workers proposed by a policy.
#[derive(Clone)]
pub struct SelectionProposal {
pub primary: Arc<Worker>,
pub backup: Option<Arc<Worker>>,
pub kind: ProposalKind,
/// Workers still eligible for fallback after filtering.
/// Optional pressure guard settings for this pair.
/// Applied only when both workers have complete, fresh native monitor data.
pub guard_hints: GuardHints,
/// Workers available for fallback after eligibility filtering.
pub eligible_workers: Option<Vec<Arc<Worker>>>,
}
/// A Cache-Aware Prefill candidate with `E = L - H`.
/// Cache-Aware prefill candidate where `E = L - H`.
#[derive(Clone)]
pub struct CacheCandidate {
pub worker: Arc<Worker>,
pub matched_prefix_tokens: u64,
pub uncached_tokens: u64,
/// Candidate domain.
/// Domain containing this candidate.
pub candidate_range_id: String,
/// Optional pending-Prefill limit checked with `E`.
/// Optional pending prefill limit checked against `E`.
pub max_pending_prefill_tokens: Option<u64>,
}
/// A bounded Cache-Aware candidate set.
#[derive(Clone)]
/// Bounded set of Cache-Aware candidates.
#[derive(Clone, Default)]
pub struct CacheCandidateProposal {
pub candidates: Vec<CacheCandidate>,
pub cache_switch_margin_tokens: u64,
pub enable_pressure_guard: bool,
pub pressure_abs_threshold_tokens: u64,
pub pressure_abs_threshold_ms: Option<f64>,
pub pressure_rel_threshold: f64,
}
/// A Prefill policy result: a pair or Cache-Aware candidates.
/// Prefill proposal returned as either a pair or a Cache-Aware candidate set.
#[derive(Clone)]
pub enum PrefillProposal {
Pair(SelectionProposal),
@@ -328,7 +356,7 @@ pub enum PrefillProposal {
}
impl PrefillProposal {
/// Applies EligibilityFilter results to either proposal form.
/// Applies eligibility filtering to either proposal form.
pub fn with_eligible_workers(self, workers: Vec<Arc<Worker>>) -> Self {
match self {
Self::Pair(proposal) => Self::Pair(proposal.with_eligible_workers(workers)),
@@ -351,6 +379,7 @@ impl SelectionProposal {
primary,
backup: None,
kind: ProposalKind::Generic,
guard_hints: GuardHints::default(),
eligible_workers: None,
}
}
@@ -361,6 +390,7 @@ impl SelectionProposal {
primary,
backup: Some(backup),
kind: ProposalKind::PowerOfTwo,
guard_hints: GuardHints::default(),
eligible_workers: None,
}
}
@@ -370,13 +400,18 @@ impl SelectionProposal {
self
}
pub fn with_guard_hints(mut self, guard_hints: GuardHints) -> Self {
self.guard_hints = guard_hints;
self
}
pub fn with_eligible_workers(mut self, workers: Vec<Arc<Worker>>) -> Self {
self.eligible_workers = Some(workers);
self
}
}
/// The source of a primary/backup proposal.
/// Source of a primary/backup proposal.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProposalKind {
Generic,
@@ -386,6 +421,26 @@ pub enum ProposalKind {
Score,
}
/// Optional guard settings for a pair proposal.
#[derive(Debug, Clone)]
pub struct GuardHints {
pub enable_pressure_guard: bool,
pub pressure_abs_threshold_tokens: u64,
pub pressure_abs_threshold_ms: Option<f64>,
pub pressure_rel_threshold: f64,
}
impl Default for GuardHints {
fn default() -> Self {
Self {
enable_pressure_guard: false,
pressure_abs_threshold_tokens: 0,
pressure_abs_threshold_ms: None,
pressure_rel_threshold: 1.0,
}
}
}
pub trait Policy: Send + Sync + std::fmt::Debug {
fn select(&self, workers: &[Arc<Worker>], ctx: &SelectionContext<'_>) -> Option<Arc<Worker>>;
@@ -426,7 +481,12 @@ pub trait Policy: Send + Sync + std::fmt::Debug {
self.uses_shared_prefill_admission()
}
/// Whether this policy resolves an affinity primary within the candidate range.
/// Whether in-flight requests must be timestamped for load correction.
fn needs_dispatch_timestamps(&self) -> bool {
false
}
/// Whether this policy resolves an affinity primary within its candidate range.
fn is_bucket_affinity_policy(&self) -> bool {
false
}
@@ -499,20 +559,23 @@ mod tests {
resolve_cache_candidates, resolve_prefill, CandidateRange, DecisionReason, FreshLoadLookup,
};
use crate::policies::cache_aware::CacheAwarePolicy;
use crate::policies::engine_load::{EngineLoadSnapshot, EngineWorkerLoad};
use crate::policies::engine_load::{EngineLoadSnapshot, NativeCacheWorkerLoad};
use crate::policies::power_of_two::PowerOfTwoChoicesPolicy;
use crate::policies::round_robin::RoundRobinPolicy;
use crate::policies::session_aware::SessionAwarePolicy;
use std::collections::HashMap;
use std::time::Instant;
/// Aggregated `LoadStat` values used only by policy tests.
/// #34608 `LoadStat` aggregate used by policy tests.
#[derive(Clone, Default)]
struct TestEngineLoad {
num_running_reqs: u64,
num_waiting_reqs: u64,
num_tokens: u64,
max_total_num_tokens: u64,
num_waiting_uncached_tokens: Option<u64>,
num_total_tokens: Option<u64>,
max_running_requests: Option<u64>,
}
fn worker(id: &str) -> Arc<Worker> {
@@ -601,6 +664,7 @@ mod tests {
max_pending_prefill_tokens: None,
}],
cache_switch_margin_tokens: 8,
..Default::default()
};
assert_eq!(proposal.candidates[0].worker.id, hot.id);
@@ -1118,18 +1182,27 @@ mod tests {
}
fn snapshot(entries: &[(&Arc<Worker>, TestEngineLoad)]) -> EngineLoadSnapshot {
EngineLoadSnapshot::from_workers(
EngineLoadSnapshot::from_native_cache_workers(
1,
entries
.iter()
.map(|(worker, aggregate)| {
(
worker.url.clone(),
EngineWorkerLoad {
NativeCacheWorkerLoad {
num_running_reqs: aggregate.num_running_reqs,
num_waiting_reqs: aggregate.num_waiting_reqs,
num_tokens: aggregate.num_tokens,
num_waiting_uncached_tokens: aggregate
.num_waiting_uncached_tokens
.unwrap_or(aggregate.num_waiting_reqs),
num_used_tokens: aggregate.num_tokens,
num_total_tokens: aggregate
.num_total_tokens
.unwrap_or(aggregate.num_tokens),
max_total_num_tokens: aggregate.max_total_num_tokens,
max_running_requests: aggregate.max_running_requests.unwrap_or(64),
prefill_throughput_tokens_per_s: None,
estimated_prefill_queue_ms: None,
captured_at: Instant::now(),
},
)
@@ -1202,6 +1275,7 @@ mod tests {
cache_candidate(&winner, 70, 30, None),
],
cache_switch_margin_tokens: 16,
..Default::default()
};
let loads = snapshot(&[
(
@@ -1223,6 +1297,7 @@ mod tests {
]);
let decision = resolve_cache_candidates(&proposal, 100, &loads)
.decision
.expect("a later admitted cache match must survive");
assert_eq!(decision.selected.id, winner.id);
@@ -1243,6 +1318,7 @@ mod tests {
cache_candidate(&final_winner, 80, 20, None),
],
cache_switch_margin_tokens: 0,
..Default::default()
};
let loads = snapshot(&[
(
@@ -1269,6 +1345,7 @@ mod tests {
]);
let decision = resolve_cache_candidates(&proposal, 100, &loads)
.decision
.expect("all admitted candidates must participate in the tournament");
assert_eq!(decision.selected.id, final_winner.id);
@@ -1282,6 +1359,7 @@ mod tests {
let proposal = CacheCandidateProposal {
candidates: vec![cache_candidate(&candidate, 80, 20, Some(30))],
cache_switch_margin_tokens: 16,
..Default::default()
};
let pending_allows = snapshot(&[(
&candidate,
@@ -1292,7 +1370,9 @@ mod tests {
},
)]);
assert!(
resolve_cache_candidates(&proposal, 100, &pending_allows).is_some(),
resolve_cache_candidates(&proposal, 100, &pending_allows)
.decision
.is_some(),
"pending admission must project E=20, not L=100"
);
@@ -1306,7 +1386,9 @@ mod tests {
},
)]);
assert!(
resolve_cache_candidates(&proposal, 100, &kv_rejects).is_none(),
resolve_cache_candidates(&proposal, 100, &kv_rejects)
.decision
.is_none(),
"KV safety must conservatively project the complete input L=100"
);
}
@@ -1321,6 +1403,7 @@ mod tests {
cache_candidate(&idle, 80, 20, None),
],
cache_switch_margin_tokens: 32,
..Default::default()
};
let loads = snapshot(&[
(
@@ -1341,7 +1424,9 @@ mod tests {
),
]);
let decision = resolve_cache_candidates(&proposal, 100, &loads).unwrap();
let decision = resolve_cache_candidates(&proposal, 100, &loads)
.decision
.unwrap();
assert_eq!(decision.selected.id, congested.id);
}
@@ -1355,6 +1440,7 @@ mod tests {
cache_candidate(&idle, 20, 80, None),
],
cache_switch_margin_tokens: 32,
..Default::default()
};
let loads = snapshot(&[
(
@@ -1375,7 +1461,9 @@ mod tests {
),
]);
let decision = resolve_cache_candidates(&proposal, 100, &loads).unwrap();
let decision = resolve_cache_candidates(&proposal, 100, &loads)
.decision
.unwrap();
assert_eq!(
decision.selected.id, hot.id,
"pressure may break a near tie, but must not erase a material cache-work gain"
@@ -1397,6 +1485,7 @@ mod tests {
cache_candidate(&beyond_margin, 60, 40, None),
],
cache_switch_margin_tokens: 32,
..Default::default()
};
let loads = snapshot(&[
(
@@ -1425,7 +1514,9 @@ mod tests {
),
]);
let decision = resolve_cache_candidates(&proposal, 100, &loads).unwrap();
let decision = resolve_cache_candidates(&proposal, 100, &loads)
.decision
.unwrap();
assert_eq!(
decision.selected.id, best_work.id,
"without a unit-compatible token-pressure signal, cache work remains authoritative"
@@ -0,0 +1,62 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use super::ExternalPrefixSignal;
use crate::policies::kv_events::{
compute_block_hashes, compute_block_hashes_bigram, BlockSizeOracle, HashTree,
};
use sgl_kv_indexer::{PrefixMatch, PrefixOutcome};
use std::collections::BTreeMap;
use std::sync::Arc;
#[derive(Clone, Debug)]
pub struct RadixTreePrefixProvider {
tree: Arc<HashTree>,
block_size_oracle: Arc<BlockSizeOracle>,
}
impl RadixTreePrefixProvider {
pub fn new(tree: Arc<HashTree>, block_size_oracle: Arc<BlockSizeOracle>) -> Self {
Self {
tree,
block_size_oracle,
}
}
pub fn match_request_tokens(&self, tokens: &[u32]) -> Option<ExternalPrefixSignal> {
let block_size = self.block_size_oracle.get()?;
let hashes = if self.block_size_oracle.is_bigram() {
compute_block_hashes_bigram(tokens, block_size as usize)
} else {
compute_block_hashes(tokens, block_size as usize)
};
if hashes.is_empty() {
return None;
}
let mut depth_by_url = BTreeMap::<String, u32>::new();
for (worker, depth) in self.tree.prefix_depths(None, &hashes) {
let depth = u32::try_from(depth).unwrap_or(u32::MAX);
depth_by_url
.entry(worker.url)
.and_modify(|current| *current = (*current).max(depth))
.or_insert(depth);
}
let best_prefix_blocks = depth_by_url.values().copied().max()?;
let matches = depth_by_url
.into_iter()
.map(|(address, matched_prefix_blocks)| PrefixMatch {
worker_id: address.clone(),
address,
matched_prefix_blocks,
})
.collect();
Some(ExternalPrefixSignal {
outcome: PrefixOutcome::Matched {
matches,
best_prefix_blocks,
},
query_blocks: hashes.len(),
})
}
}
@@ -228,7 +228,7 @@ impl PdPoolResolver {
///
/// The current `Policy` trait carries `(workers, ctx)`; adding an
/// `affinity_hint` argument would touch every policy implementation
/// (`round_robin`, `random`, `power_of_two`, `cache_aware_zmq`).
/// (`round_robin`, `random`, `power_of_two`, `cache_aware`).
/// Affinity is a PD-routing concern — orthogonal to the in-pool
/// scoring the trait abstracts — so keeping it as a sibling helper
/// keeps the trait's responsibility narrow.
@@ -57,6 +57,11 @@ pub trait ScoringPolicy: Send + Sync + std::fmt::Debug {
false
}
/// Whether scoring corrects Engine Load with recent dispatch timestamps.
fn needs_dispatch_timestamps(&self) -> bool {
false
}
/// Optional eligibility view for policies that provide both signals.
fn as_filter(&self) -> Option<&dyn EligibilityFilter> {
None
@@ -143,6 +148,10 @@ impl<T: ScoringPolicy> Policy for T {
ScoringPolicy::needs_load_snapshot(self)
}
fn needs_dispatch_timestamps(&self) -> bool {
ScoringPolicy::needs_dispatch_timestamps(self)
}
fn as_scoring(&self) -> Option<&dyn ScoringPolicy> {
Some(self)
}
@@ -281,6 +290,14 @@ impl Policy for Pipeline {
self.inner.needs_load_snapshot() || self.filters.iter().any(|p| p.needs_load_snapshot())
}
fn needs_dispatch_timestamps(&self) -> bool {
self.inner.needs_dispatch_timestamps()
|| self
.filters
.iter()
.any(|policy| policy.needs_dispatch_timestamps())
}
fn commit_prefill_selection(
&self,
ctx: &SelectionContext<'_>,
@@ -340,6 +357,10 @@ impl Policy for ScorePolicy {
self.inner.needs_request_tokens()
}
fn needs_dispatch_timestamps(&self) -> bool {
self.inner.needs_dispatch_timestamps()
}
fn attach_metrics(&self, metrics: Arc<crate::server::metrics::MetricsRegistry>) {
self.inner.attach_metrics(metrics);
}
@@ -373,6 +394,12 @@ impl ScoringPolicy for FusedScorePolicy {
.iter()
.any(|(policy, _)| policy.needs_load_snapshot())
}
fn needs_dispatch_timestamps(&self) -> bool {
self.terms
.iter()
.any(|(policy, _)| policy.needs_dispatch_timestamps())
}
}
/// Owned boxes as the borrowed views [`admit`] consumes. Shared by the tests
@@ -390,7 +417,8 @@ mod tests {
use crate::config::AffinityConfig;
use crate::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use crate::policies::admission::{resolve_prefill, CandidateRange};
use crate::policies::engine_load::{EngineLoadSnapshot, EngineWorkerLoad};
use crate::policies::engine_load::{EngineLoadSnapshot, NativeCacheWorkerLoad};
use crate::policies::load_based::LoadBasedPolicy;
use crate::policies::power_of_two::PowerOfTwoChoicesPolicy;
use crate::policies::round_robin::RoundRobinPolicy;
use crate::policies::session_aware::SessionAwarePolicy;
@@ -412,18 +440,23 @@ mod tests {
}
fn snapshot(entries: &[(&Arc<Worker>, u64, u64, u64, u64)]) -> EngineLoadSnapshot {
EngineLoadSnapshot::from_workers(
EngineLoadSnapshot::from_native_cache_workers(
1,
entries
.iter()
.map(|(worker, running, waiting, used, capacity)| {
(
worker.url.clone(),
EngineWorkerLoad {
NativeCacheWorkerLoad {
num_running_reqs: *running,
num_waiting_reqs: *waiting,
num_tokens: *used,
num_waiting_uncached_tokens: *waiting,
num_used_tokens: *used,
num_total_tokens: *used,
max_total_num_tokens: *capacity,
max_running_requests: 64,
prefill_throughput_tokens_per_s: None,
estimated_prefill_queue_ms: None,
captured_at: Instant::now(),
},
)
@@ -606,22 +639,36 @@ mod tests {
let plain = FusedScorePolicy::new(vec![term(by(1.0), None)]).unwrap();
assert!(!Policy::needs_load_snapshot(&plain));
assert!(!Policy::needs_dispatch_timestamps(&plain));
let fused =
FusedScorePolicy::new(vec![term(by(1.0), None), term(LoadHungry, None)]).unwrap();
assert!(Policy::needs_load_snapshot(&fused));
assert!(!Policy::needs_dispatch_timestamps(&fused));
let load_fused = FusedScorePolicy::new(vec![
term(by(1.0), None),
term(LoadBasedPolicy::new(), None),
])
.unwrap();
assert!(Policy::needs_dispatch_timestamps(&load_fused));
let pipeline = Pipeline::new(
vec![Arc::new(Keep(vec!["a"], OnEmpty::Abstain))],
Arc::new(fused),
Arc::new(load_fused),
)
.unwrap();
assert!(pipeline.needs_load_snapshot());
assert!(pipeline.needs_dispatch_timestamps());
let score = ScorePolicy::new(Arc::new(by(1.0)));
assert!(
score.needs_load_snapshot(),
"shared admission requires a snapshot"
);
assert!(!score.needs_dispatch_timestamps());
let load_score = ScorePolicy::new(Arc::new(LoadBasedPolicy::new()));
assert!(load_score.needs_dispatch_timestamps());
}
#[test]
@@ -8,7 +8,7 @@ use crate::discovery::WorkerId;
use crate::policies::active_load::{spawn_sweeper, Clock, JanitorHandle, SystemTimeClock};
use crate::policies::admission::compare_prefill_pressure;
use crate::policies::power_of_two::PowerOfTwoChoicesPolicy;
use crate::policies::{Policy, ProposalKind, SelectionContext, SelectionProposal};
use crate::policies::{GuardHints, Policy, ProposalKind, SelectionContext, SelectionProposal};
use crate::workers::Worker;
use dashmap::DashMap;
use rand::Rng;
@@ -139,7 +139,15 @@ impl SessionAwarePolicy {
Some(backup) => SelectionProposal::with_backup(primary, backup),
None => SelectionProposal::primary(primary),
};
proposal.with_kind(ProposalKind::SessionAffinity)
proposal
.with_kind(ProposalKind::SessionAffinity)
.with_guard_hints(GuardHints {
enable_pressure_guard: self.config.pressure_guard
&& self.config.mode == crate::config::AffinityMode::Soft,
pressure_abs_threshold_tokens: self.config.pressure_abs_threshold_tokens,
pressure_abs_threshold_ms: self.config.pressure_abs_threshold_ms,
pressure_rel_threshold: self.config.pressure_rel_threshold,
})
}
}
@@ -216,6 +216,10 @@ impl Policy for StickyPolicy {
fn needs_load_snapshot(&self) -> bool {
self.fallback.needs_load_snapshot()
}
fn needs_dispatch_timestamps(&self) -> bool {
self.fallback.needs_dispatch_timestamps()
}
}
impl std::fmt::Debug for StickyPolicy {
@@ -241,6 +245,14 @@ mod tests {
Arc::new(crate::policies::power_of_two::PowerOfTwoChoicesPolicy::new()),
);
assert!(policy.needs_load_snapshot());
assert!(!policy.needs_dispatch_timestamps());
let load_policy = StickyPolicy::new(
Duration::from_secs(60),
Duration::from_secs(10),
Arc::new(crate::policies::load_based::LoadBasedPolicy::new()),
);
assert!(load_policy.needs_dispatch_timestamps());
}
use crate::policies::round_robin::RoundRobinPolicy;
@@ -4,8 +4,10 @@
use crate::config::Config;
use crate::policies::active_load::ActiveLoadRegistry;
use crate::policies::buckets::BucketSelector;
use crate::policies::engine_load::EngineLoadTable;
use crate::policies::kv_events::BlockSizeOracle;
use crate::policies::prefix_provider::RadixTreePrefixProvider;
use crate::policies::PolicyRegistry;
use crate::proxy::Proxy;
use crate::server::metrics::MetricsRegistry;
@@ -20,17 +22,19 @@ pub struct AppContext {
pub proxy: Arc<Proxy>,
pub registry: Arc<WorkerRegistry>,
pub policies: Arc<PolicyRegistry>,
/// Converts static Bucket configuration into request candidate domains.
pub bucket_selector: Arc<BucketSelector>,
/// Per-worker active-load bookkeeping shared by the proxy, policies,
/// timeout janitor, and metrics.
pub active_load: Arc<ActiveLoadRegistry>,
/// Lightweight Prometheus-format metrics registry served via
/// `/metrics`. Shared with the chat handler (requests_total),
/// cache-aware-zmq policy (overlap_blocks), active-load registry
/// (active_load gauge + stale_requests_total), and PD dispatch.
/// active-load registry, policy-specific counters, and PD dispatch.
pub metrics: Arc<MetricsRegistry>,
/// Shared Engine LoadStat table; ingress captures one immutable snapshot per request.
pub engine_load: Arc<EngineLoadTable>,
pub prefix_index: Option<Arc<sgl_kv_indexer::GrpcPrefixIndex>>,
pub prefix_index: Option<Arc<dyn sgl_kv_indexer::PrefixIndex>>,
pub radix_tree_prefix_provider: Option<RadixTreePrefixProvider>,
pub block_size_oracle: Arc<BlockSizeOracle>,
ready: AtomicBool,
}
@@ -71,20 +75,21 @@ impl AppContext {
// Without this, the metric is permanently 0 in production even
// though the chat handler is faithfully calling `register`.
active_load.attach_metrics(Arc::clone(&metrics));
// Same rationale for the cache-aware-zmq policy's
// `sgl_router_overlap_blocks`: the metrics registry is built here,
// after the policy registry, so inject it now. No-op for policies
// that don't emit metrics.
// The metrics registry is built after the policy registry, so attach
// it here for policies that emit their own counters.
policies.attach_metrics(Arc::clone(&metrics));
let bucket_selector = Arc::new(BucketSelector::new(config.model.bucket_config.clone()));
Self {
config,
tokenizers,
proxy,
registry,
policies,
bucket_selector,
active_load,
metrics,
prefix_index: None,
radix_tree_prefix_provider: None,
block_size_oracle: BlockSizeOracle::new(),
engine_load: EngineLoadTable::new(),
ready: AtomicBool::new(false),
@@ -114,6 +119,8 @@ impl AppContext {
id: "stub-model".into(),
tokenizer_path: "stub".into(),
policy: crate::config::PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -133,9 +140,11 @@ impl AppContext {
proxy: Arc::new(Proxy::new(std::time::Duration::from_secs(60)).expect("stub proxy")),
registry: Arc::new(WorkerRegistry::default()),
policies: Arc::new(PolicyRegistry::default()),
bucket_selector: Arc::new(BucketSelector::new(None)),
active_load: ActiveLoadRegistry::with_defaults(),
metrics: MetricsRegistry::new(),
prefix_index: None,
radix_tree_prefix_provider: None,
block_size_oracle: BlockSizeOracle::new(),
engine_load: EngineLoadTable::new(),
ready: AtomicBool::new(false),
+114 -74
View File
@@ -23,7 +23,6 @@
//! | `sgl_router_worker_requests_total` | Counter | `worker_url`, `model_id`, `mode`, `outcome` |
//! | `sgl_router_request_duration_seconds` | Histogram | `model_id` |
//! | `sgl_router_ttft_seconds` | Histogram | `model_id` |
//! | `sgl_router_overlap_blocks` | Histogram | `model_id` |
//! | `sgl_router_active_load` | Gauge | `worker_url`, `kind` |
//! | `sgl_router_workers` | Gauge | `mode` |
//! | `sgl_router_worker_health` | Gauge | `worker_url` |
@@ -34,6 +33,11 @@
//! | `sgl_router_sticky_total` | Counter | `outcome` |
//! | `sgl_router_policy_decisions_total` | Counter | `policy`, `reason` |
//! | `sgl_router_policy_selection_failures_total` | Counter | `policy`, `reason` |
//! | `sgl_router_cache_admission_evaluated_total` | Counter | — |
//! | `sgl_router_cache_admission_rejected_total` | Counter | — |
//! | `sgl_router_cache_pressure_guard_compared_total` | Counter | — |
//! | `sgl_router_cache_pressure_guard_override_total` | Counter | — |
//! | `sgl_router_cache_monitor_decisions_total` | Counter | `source` |
//! | `sgl_router_ingress_tokenize_errors_total` | Counter | `model_id` |
//!
//! The four `sgl_router_worker*` gauges and `sgl_router_workers` are sampled
@@ -50,16 +54,6 @@ use std::collections::HashMap;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::sync::Arc;
/// Histogram bucket upper bounds for `sgl_router_overlap_blocks`. Blocks are
/// 3264 tokens each, and the `MAX_CHAT_BODY_BYTES` cap bounds context length —
/// putting the practical ceiling for a maximum-length context in the low tens
/// of thousands of blocks. The ladder spans 0 → ~8k blocks at the resolution
/// worth charting; the `+Inf` bucket catches the longer-context tail beyond
/// 8000.
const OVERLAP_BLOCKS_BUCKETS: &[f64] = &[
0.0, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0, 128.0, 256.0, 512.0, 1000.0, 2000.0, 4000.0, 8000.0,
];
/// Histogram bucket upper bounds (seconds) for
/// `sgl_router_request_duration_seconds`. Standard latency ladder spanning
/// 5 ms → 30 s; the `+Inf` bucket catches anything slower (a request that
@@ -239,13 +233,17 @@ pub struct MetricsRegistry {
// on `worker_requests_total` / the worker gauges instead.
request_duration: Mutex<HashMap<String, Histogram>>,
ttft_seconds: Mutex<HashMap<String, Histogram>>,
overlap_blocks: Mutex<HashMap<String, Histogram>>,
active_load: Mutex<HashMap<ActiveLoadKey, Arc<AtomicI64>>>,
stale_requests_total: Mutex<HashMap<&'static str, Arc<AtomicU64>>>,
decode_affinity_total: Mutex<HashMap<&'static str, Arc<AtomicU64>>>,
sticky_total: Mutex<HashMap<&'static str, Arc<AtomicU64>>>,
policy_decisions_total: Mutex<HashMap<PolicyDecisionKey, Arc<AtomicU64>>>,
policy_selection_failures_total: Mutex<HashMap<PolicyDecisionKey, Arc<AtomicU64>>>,
cache_admission_evaluated_total: AtomicU64,
cache_admission_rejected_total: AtomicU64,
cache_pressure_guard_compared_total: AtomicU64,
cache_pressure_guard_override_total: AtomicU64,
cache_monitor_decisions_total: Mutex<HashMap<&'static str, Arc<AtomicU64>>>,
ingress_tokenize_errors_total: Mutex<HashMap<String, Arc<AtomicU64>>>,
}
@@ -304,10 +302,8 @@ struct PolicyDecisionKey {
#[derive(Debug)]
struct Histogram {
/// Bucket upper bounds this histogram observes against (e.g.
/// [`OVERLAP_BLOCKS_BUCKETS`] or [`REQUEST_DURATION_BUCKETS`]). Held
/// per-instance so a single `Histogram` type backs metrics with
/// different bucket ladders.
/// Bucket upper bounds this histogram observes against. Held per-instance
/// so a single `Histogram` type backs metrics with different bucket ladders.
bounds: &'static [f64],
/// One counter per boundary in `bounds`, plus one for `+Inf`. Buckets
/// are cumulative on render but stored as non-cumulative counts here.
@@ -394,15 +390,6 @@ impl MetricsRegistry {
counter.fetch_add(1, Ordering::Relaxed);
}
/// Observe an overlap-blocks count for `sgl_router_overlap_blocks`.
pub fn observe_overlap_blocks(&self, model_id: &str, blocks: u64) {
let mut guard = self.overlap_blocks.lock();
let hist = guard
.entry(model_id.to_owned())
.or_insert_with(|| Histogram::new(OVERLAP_BLOCKS_BUCKETS));
hist.observe(blocks as f64);
}
/// Observe end-to-end request latency (seconds) for
/// `sgl_router_request_duration_seconds`. Recorded once the upstream
/// outcome is known, regardless of success or error — a slow error is
@@ -542,6 +529,38 @@ impl MetricsRegistry {
counter.fetch_add(1, Ordering::Relaxed);
}
/// Cache-Aware candidates evaluated by hard admission.
pub fn record_cache_admission_evaluations(&self, count: u64) {
self.cache_admission_evaluated_total
.fetch_add(count, Ordering::Relaxed);
}
/// Cache-Aware candidates rejected by hard admission.
pub fn record_cache_admission_rejections(&self, count: u64) {
self.cache_admission_rejected_total
.fetch_add(count, Ordering::Relaxed);
}
/// Pressure-guard pairs compared and overridden with complete monitor data.
pub fn record_cache_pressure_guard(&self, compared: u64, overrides: u64) {
self.cache_pressure_guard_compared_total
.fetch_add(compared, Ordering::Relaxed);
self.cache_pressure_guard_override_total
.fetch_add(overrides, Ordering::Relaxed);
}
/// Load source used for a Cache-Aware decision. Benchmarks reject
/// `router_local` results to verify that monitor data affected selection.
pub fn record_cache_monitor_decision(&self, source: &'static str) {
let mut guard = self.cache_monitor_decisions_total.lock();
let counter = guard
.entry(source)
.or_insert_with(|| Arc::new(AtomicU64::new(0)))
.clone();
drop(guard);
counter.fetch_add(1, Ordering::Relaxed);
}
/// Bump `sgl_router_ingress_tokenize_errors_total{model_id}`.
///
/// Recorded ONLY when the tokenization offload SHOULD have fired but the
@@ -693,21 +712,6 @@ impl MetricsRegistry {
}
drop(guard);
// overlap_blocks histogram
out.push_str(
"# HELP sgl_router_overlap_blocks Overlap-block count observed at cache-aware-zmq policy selection.\n",
);
out.push_str("# TYPE sgl_router_overlap_blocks histogram\n");
let guard = self.overlap_blocks.lock();
let mut models: Vec<&String> = guard.keys().collect();
models.sort();
for model_id in models {
let hist = guard.get(model_id).unwrap();
let label_body = format!("model_id=\"{}\"", escape_label(model_id));
render_histogram(&mut out, "sgl_router_overlap_blocks", &label_body, hist);
}
drop(guard);
// active_load gauge
out.push_str(
"# HELP sgl_router_active_load Per-worker active load (prefill_tokens or decode_blocks).\n",
@@ -888,6 +892,58 @@ impl MetricsRegistry {
}
drop(guard);
out.push_str(
"# HELP sgl_router_cache_admission_evaluated_total Cache-Aware candidates evaluated by hard admission.\n",
);
out.push_str("# TYPE sgl_router_cache_admission_evaluated_total counter\n");
out.push_str(&format!(
"sgl_router_cache_admission_evaluated_total {}\n",
self.cache_admission_evaluated_total.load(Ordering::Relaxed),
));
out.push_str(
"# HELP sgl_router_cache_admission_rejected_total Cache-Aware candidates rejected by hard admission.\n",
);
out.push_str("# TYPE sgl_router_cache_admission_rejected_total counter\n");
out.push_str(&format!(
"sgl_router_cache_admission_rejected_total {}\n",
self.cache_admission_rejected_total.load(Ordering::Relaxed),
));
out.push_str(
"# HELP sgl_router_cache_pressure_guard_compared_total Complete fresh Cache-Aware candidate pairs evaluated by the pressure guard.\n",
);
out.push_str("# TYPE sgl_router_cache_pressure_guard_compared_total counter\n");
out.push_str(&format!(
"sgl_router_cache_pressure_guard_compared_total {}\n",
self.cache_pressure_guard_compared_total
.load(Ordering::Relaxed),
));
out.push_str(
"# HELP sgl_router_cache_pressure_guard_override_total Pressure-guard comparisons whose outcome differs from cache/work ordering without the guard.\n",
);
out.push_str("# TYPE sgl_router_cache_pressure_guard_override_total counter\n");
out.push_str(&format!(
"sgl_router_cache_pressure_guard_override_total {}\n",
self.cache_pressure_guard_override_total
.load(Ordering::Relaxed),
));
out.push_str(
"# HELP sgl_router_cache_monitor_decisions_total Cache-Aware candidate resolutions by actual load source.\n",
);
out.push_str("# TYPE sgl_router_cache_monitor_decisions_total counter\n");
let guard = self.cache_monitor_decisions_total.lock();
let mut entries: Vec<(&&str, u64)> = guard
.iter()
.map(|(source, value)| (source, value.load(Ordering::Relaxed)))
.collect();
entries.sort_by_key(|entry| *entry.0);
for (source, value) in entries {
out.push_str(&format!(
"sgl_router_cache_monitor_decisions_total{{source=\"{}\"}} {}\n",
source, value,
));
}
drop(guard);
// ingress_tokenize_errors_total
out.push_str(
"# HELP sgl_router_ingress_tokenize_errors_total Chat requests on a chat-encoder model whose ingress tokenization failed, silently falling back to engine-side tokenization (the input_ids offload was defeated).\n",
@@ -964,7 +1020,6 @@ mod tests {
assert!(out.contains("# TYPE sgl_router_request_duration_seconds histogram"));
assert!(out.contains("# TYPE sgl_router_ttft_seconds histogram"));
assert!(out.contains("# TYPE sgl_router_responses_total counter"));
assert!(out.contains("# TYPE sgl_router_overlap_blocks histogram"));
assert!(out.contains("# TYPE sgl_router_active_load gauge"));
assert!(out.contains("# TYPE sgl_router_workers gauge"));
assert!(out.contains("# TYPE sgl_router_worker_health gauge"));
@@ -1200,28 +1255,6 @@ mod tests {
);
}
#[test]
fn observe_overlap_blocks_writes_buckets_and_count() {
let reg = MetricsRegistry::new();
reg.observe_overlap_blocks("tiny", 3);
reg.observe_overlap_blocks("tiny", 9);
reg.observe_overlap_blocks("tiny", 50);
let out = reg.render();
// 3 observations -> count=3, sum=62
assert!(out.contains(r#"sgl_router_overlap_blocks_count{model_id="tiny"} 3"#));
assert!(out.contains(r#"sgl_router_overlap_blocks_sum{model_id="tiny"} 62"#));
// The le=64 bucket is cumulative: 3 is <=4, 9 is <=16, 50 is <=64.
assert!(
out.contains(r#"sgl_router_overlap_blocks_bucket{model_id="tiny",le="64"} 3"#),
"bucket le=64 should be 3 (cumulative); got:\n{out}",
);
// The le=4 bucket should include only the 3.
assert!(
out.contains(r#"sgl_router_overlap_blocks_bucket{model_id="tiny",le="4"} 1"#),
"bucket le=4 should be 1; got:\n{out}",
);
}
#[test]
fn set_active_load_gauge_overwrites() {
let reg = MetricsRegistry::new();
@@ -1320,6 +1353,24 @@ mod tests {
));
}
#[test]
fn cache_monitor_and_guard_counters_are_exposed() {
let reg = MetricsRegistry::new();
reg.record_cache_monitor_decision("estimated_prefill_queue_ms");
reg.record_cache_admission_evaluations(3);
reg.record_cache_admission_rejections(2);
reg.record_cache_pressure_guard(3, 1);
let out = reg.render();
assert!(out.contains(
r#"sgl_router_cache_monitor_decisions_total{source="estimated_prefill_queue_ms"} 1"#
));
assert!(out.contains("sgl_router_cache_admission_evaluated_total 3"));
assert!(out.contains("sgl_router_cache_admission_rejected_total 2"));
assert!(out.contains("sgl_router_cache_pressure_guard_compared_total 3"));
assert!(out.contains("sgl_router_cache_pressure_guard_override_total 1"));
}
#[test]
fn ingress_tokenize_error_counter_increments_per_model() {
let reg = MetricsRegistry::new();
@@ -1369,15 +1420,4 @@ mod tests {
"render did not escape backslash; got:\n{out}",
);
}
#[test]
fn histogram_plus_inf_bucket_catches_overflow() {
let reg = MetricsRegistry::new();
// 8001 is just above the last finite bucket (8000); it should land
// in +Inf only.
reg.observe_overlap_blocks("m", 8001);
let out = reg.render();
assert!(out.contains(r#"sgl_router_overlap_blocks_bucket{model_id="m",le="8000"} 0"#));
assert!(out.contains(r#"sgl_router_overlap_blocks_bucket{model_id="m",le="+Inf"} 1"#));
}
}
+613 -108
View File
@@ -1,8 +1,16 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use crate::config::{PolicyKind, SessionAffinityMode};
use crate::discovery::{ModelId, WorkerMode};
use crate::policies::admission::{resolve_cache_candidates, resolve_prefill, CandidateRange};
use crate::policies::admission::{
resolve_cache_candidates, resolve_decode, resolve_prefill, resolve_prefill_admitted,
CandidateDomain, CandidateRange, DecisionReason,
};
use crate::policies::buckets::BucketRequest;
use crate::policies::decode::{
build_decode_policy, resolve_decode_with_capacity_fallback, DecodeSelectionContext,
};
use crate::policies::kv_events::{compute_block_hashes, compute_block_hashes_bigram};
use crate::policies::registry::{PdPoolResolver, PdResolveError};
use crate::policies::{
@@ -22,12 +30,12 @@ use axum::http::{HeaderMap, HeaderName, HeaderValue, Response};
use bytes::Bytes;
use serde::de::IgnoredAny;
use serde::Deserialize;
use sgl_kv_indexer::PrefixIndex;
use std::cell::Cell;
use std::collections::HashMap;
use std::sync::Arc;
/// Observability header carrying the decode-pool URL selected via host
/// affinity for a PD-disaggregated request. The router fans the
/// Observability header carrying the final decode-pool URL for a
/// PD-disaggregated request. The router fans the
/// bootstrap-injected request body to BOTH the prefill and the decode
/// worker concurrently; this header lets the prefill log the chosen
/// peer, and is mirrored onto the response so sidecars / tests can
@@ -35,6 +43,10 @@ use std::sync::Arc;
/// prefix matches `x-sgl-router-error-code` so router-emitted metadata
/// stays grouped.
const X_SGL_DECODE_URL: HeaderName = HeaderName::from_static("x-sgl-decode-url");
/// Optional caller requirement consumed only when a static P Bucket config is enabled.
const X_SGL_TTFT_SLO_MS: HeaderName = HeaderName::from_static("x-sgl-ttft-slo-ms");
/// Optional caller TPS requirement consumed only when a static D Bucket config is enabled.
const X_SGL_TPS_SLO: HeaderName = HeaderName::from_static("x-sgl-tps-slo");
/// Coarse char-count → token-count divisor used to estimate prefill load
/// from the request body when no real tokenizer count is available. Four
@@ -46,6 +58,55 @@ const X_SGL_DECODE_URL: HeaderName = HeaderName::from_static("x-sgl-decode-url")
/// purpose.
const CHARS_PER_TOKEN_ESTIMATE: usize = 4;
/// Return the low-cardinality reason for the final Prefill decision.
fn prefill_policy_reason(
policy: PolicyKind,
proposal: ProposalKind,
decision: DecisionReason,
has_session_id: bool,
affinity_lookup_enabled: bool,
) -> &'static str {
match policy {
PolicyKind::SessionAware => match (proposal, decision) {
(ProposalKind::SessionAffinity, DecisionReason::Primary) => "session_primary",
(ProposalKind::SessionAffinity, DecisionReason::BackupPrimaryAdmission) => {
"session_admission_backup"
}
(ProposalKind::SessionAffinity, DecisionReason::BackupPressureGuard) => {
"session_pressure_backup"
}
(ProposalKind::SessionAffinity, DecisionReason::RangeFallback) => {
"session_range_fallback"
}
(_, DecisionReason::CapacityFallbackPowerOfTwo) => "capacity_fallback_power_of_two",
(_, DecisionReason::RangeFallback) => "range_fallback",
(_, _) if !affinity_lookup_enabled => "range_fallback",
(_, _) if !has_session_id => "no_session",
(ProposalKind::PowerOfTwo, _) => "assigned",
_ => "primary",
},
PolicyKind::CacheAware => match (proposal, decision) {
(_, DecisionReason::CacheCandidate)
| (ProposalKind::CacheAffinity, DecisionReason::Primary) => "cache_candidate",
(_, DecisionReason::Primary) => "no_cache_candidate",
(_, DecisionReason::BackupPrimaryAdmission) => "no_cache_candidate_admission_backup",
(_, DecisionReason::BackupPressureGuard) => "no_cache_candidate_pressure_backup",
(_, DecisionReason::RangeFallback) => "no_cache_candidate_range_fallback",
(_, DecisionReason::CapacityFallbackPowerOfTwo) => {
"no_cache_candidate_capacity_fallback_power_of_two"
}
},
_ => match decision {
DecisionReason::Primary => "primary",
DecisionReason::CacheCandidate => "cache_candidate",
DecisionReason::BackupPrimaryAdmission => "admission_backup",
DecisionReason::BackupPressureGuard => "pressure_backup",
DecisionReason::RangeFallback => "range_fallback",
DecisionReason::CapacityFallbackPowerOfTwo => "capacity_fallback_power_of_two",
},
}
}
/// Per-route body-size cap on `/v1/chat/completions`. 5 MiB accommodates a
/// long context — a ~1 M-token context tokenized as JSON fits under this —
/// while preventing a hostile client from forcing the router to
@@ -73,6 +134,24 @@ struct RequestProbe {
stream: Option<bool>,
#[serde(default)]
model: Option<String>,
/// Explicit output budget used by Decode Bucket routing.
#[serde(default)]
max_tokens: Option<u64>,
#[serde(default)]
max_completion_tokens: Option<u64>,
}
impl RequestProbe {
fn requested_max_output_tokens(&self) -> Option<u64> {
self.max_completion_tokens.or(self.max_tokens)
}
}
/// Project the peak sequence length without integer wraparound.
fn projected_decode_kv_tokens(input_tokens: u64, max_output_tokens: Option<u64>) -> u64 {
max_output_tokens.map_or(input_tokens, |output_tokens| {
input_tokens.saturating_add(output_tokens)
})
}
/// RAII guard that records `sgl_router_request_duration_seconds` when
@@ -126,6 +205,7 @@ pub async fn chat_completions(
let start = std::time::Instant::now();
let probe = parse_probe(&body)?;
let streaming = probe.stream.unwrap_or(false);
let requested_max_output_tokens = probe.requested_max_output_tokens();
let model_str = probe
.model
.ok_or_else(|| ApiError::BadRequest("missing `model` field".into()))?;
@@ -170,12 +250,18 @@ pub async fn chat_completions(
// chat encoder (`/v1/completions` / `text`), which the first gate
// alone wouldn't trigger.
//
// When neither holds, `parse_probe`'s minimal probe is enough, so we keep
// * Bucket routing also needs the prompt token count.
//
// When none holds, `parse_probe`'s minimal probe is enough, so we keep
// avoiding the full `serde_json::Value` allocation over a (up to 1 MiB)
// body. When parsed, this single value is reused for the routing
// tokenization and the outgoing-body injection below (and PD bootstrap
// injection). `parse_probe` already validated the object shape.
let want_tokens = ctx.tokenizers.has_chat_encoder(&model_str) || policy.needs_request_tokens();
let want_tokens = should_tokenize_request(
ctx.tokenizers.has_chat_encoder(&model_str),
policy.needs_request_tokens(),
ctx.bucket_selector.is_enabled(),
);
let request_value: Option<serde_json::Value> = if want_tokens {
Some(serde_json::from_slice(&body).map_err(|_| {
ApiError::BadRequest("invalid request: body must be a JSON object".into())
@@ -214,21 +300,34 @@ pub async fn chat_completions(
query_blocks,
})
}
(Some(_), _, _) => Some(ExternalPrefixSignal {
outcome: sgl_kv_indexer::PrefixOutcome::Empty,
query_blocks: 0,
}),
_ => None,
_ => ctx
.radix_tree_prefix_provider
.as_ref()
.zip(request_tokens.as_ref())
.and_then(|(provider, tokens)| provider.match_request_tokens(&tokens.ids)),
};
// Prefer exact ingress tokens; otherwise use the conservative estimate.
let prefill_load = request_tokens
.as_ref()
.map(|tokens| tokens.ids.len().max(1))
.unwrap_or_else(|| estimate_prefill_tokens(&body));
let request_input_tokens = prefill_load as u64;
let needs_load_snapshot = policy.needs_load_snapshot();
let needs_load_snapshot = policy.needs_load_snapshot()
|| workers
.iter()
.any(|worker| worker.mode() == WorkerMode::Prefill);
let load_snapshot =
needs_load_snapshot.then(|| ctx.engine_load.capture_snapshot(std::time::Instant::now()));
let needs_dispatch_timestamps = policy.needs_dispatch_timestamps();
let (ttft_slo_ms, tps_slo) = if ctx.bucket_selector.is_enabled() {
(
parse_optional_positive_u64_header(&headers, &X_SGL_TTFT_SLO_MS, "TTFT SLO")?,
parse_optional_positive_f64_header(&headers, &X_SGL_TPS_SLO, "TPS SLO")?,
)
} else {
(None, None)
};
// Sticky-session routing key. When the sticky policy is configured,
// read the routing key from the operator-chosen header into the
@@ -250,94 +349,340 @@ pub async fn chat_completions(
.and_then(|config| headers.get(config.session_id_header.as_str()))
.and_then(|value| value.to_str().ok())
.filter(|value| !value.is_empty());
let candidate_range = CandidateRange::global(&workers);
let mut selection_ctx = SelectionContext::with_routing_key(&model_id, Some(&body), routing_key)
.with_session_id(session_id)
.with_candidate_range_id(candidate_range.id)
.with_input_tokens(request_input_tokens)
.with_request_tokens(request_tokens.as_ref().map(|tokens| tokens.ids.as_slice()))
.with_external_prefix(external_prefix.as_ref());
if let Some(snapshot) = load_snapshot.as_ref() {
selection_ctx = selection_ctx.with_load_snapshot(snapshot);
}
let worker = match policy.propose_prefill(candidate_range.workers, &selection_ctx) {
Some(PrefillProposal::Pair(proposal)) if policy.uses_shared_prefill_admission() => {
let snapshot = load_snapshot
.as_ref()
.expect("shared prefill admission requires a load snapshot");
let decision =
resolve_prefill(&candidate_range, &proposal, request_input_tokens, snapshot)
.ok_or_else(|| {
policy_selection_failed(
&ctx,
&model_str,
PolicySelectionFailureReason::PrefillAdmissionExhausted,
)
})?;
policy.commit_prefill_selection(&selection_ctx, proposal.kind, &decision.selected);
decision.selected
}
Some(PrefillProposal::Pair(proposal)) => proposal.primary,
Some(PrefillProposal::CacheCandidates(proposal)) => {
let snapshot = load_snapshot
.as_ref()
.expect("cache candidate resolution requires a load snapshot");
let decision = resolve_cache_candidates(&proposal, request_input_tokens, snapshot)
.ok_or_else(|| {
policy_selection_failed(
&ctx,
&model_str,
PolicySelectionFailureReason::CacheCandidatesExhausted,
// Each Bucket retry rebuilds the proposal and reruns Admission/Guard.
let prefill_bucket_request = BucketRequest {
input_tokens: request_input_tokens,
expected_peak_sequence_tokens: None,
ttft_slo_ms,
tps_slo,
};
let configured_session_affinity_mode = ctx
.config
.model
.affinity
.as_ref()
.map(|config| config.session_affinity_mode)
.unwrap_or(SessionAffinityMode::Bucket);
// Without Bucket partitioning all modes reduce to the single global domain.
let session_affinity_mode = if ctx.bucket_selector.is_enabled() {
configured_session_affinity_mode
} else {
SessionAffinityMode::Bucket
};
let use_global_affinity_probe = ctx.bucket_selector.is_enabled()
&& policy.is_bucket_affinity_policy()
&& session_affinity_mode != SessionAffinityMode::Bucket;
let worker = {
let selection_failure_reason = Cell::new(PolicySelectionFailureReason::ProposalEmpty);
let select_prefill_in_domain = |domain: &CandidateDomain,
affinity_lookup_enabled: bool,
affinity_assignment_enabled: bool,
allow_capacity_fallback: bool|
-> Option<Arc<Worker>> {
let candidate_range = domain.prefill_range()?;
let mut selection_ctx =
SelectionContext::with_routing_key(&model_id, Some(&body), routing_key)
.with_session_id(session_id)
.with_candidate_range_id(candidate_range.id)
.with_input_tokens(request_input_tokens)
.with_request_tokens(
request_tokens.as_ref().map(|tokens| tokens.ids.as_slice()),
)
})?;
policy.commit_prefill_selection(
&selection_ctx,
ProposalKind::CacheAffinity,
&decision.selected,
);
decision.selected
}
None => {
return Err(policy_selection_failed(
&ctx,
&model_str,
PolicySelectionFailureReason::ProposalEmpty,
));
}
.with_external_prefix(external_prefix.as_ref());
if let Some(snapshot) = load_snapshot.as_ref() {
selection_ctx = selection_ctx.with_load_snapshot(snapshot);
}
let selection_ctx = if !affinity_lookup_enabled {
selection_ctx.without_affinity_lookup()
} else if !affinity_assignment_enabled {
selection_ctx.without_affinity_assignment()
} else {
selection_ctx
};
let Some(PrefillProposal::Pair(proposal)) =
policy.propose_prefill(candidate_range.workers, &selection_ctx)
else {
// Domain retries are ordinary pair proposals.
return None;
};
if policy.uses_shared_prefill_admission() {
let snapshot = load_snapshot
.as_ref()
.expect("shared prefill admission requires a load snapshot");
let decision = if allow_capacity_fallback {
resolve_prefill(&candidate_range, &proposal, request_input_tokens, snapshot)
} else {
resolve_prefill_admitted(
&candidate_range,
&proposal,
request_input_tokens,
snapshot,
)
};
let Some(decision) = decision else {
selection_failure_reason
.set(PolicySelectionFailureReason::PrefillAdmissionExhausted);
return None;
};
let reason = prefill_policy_reason(
ctx.config.model.policy,
proposal.kind,
decision.reason,
session_id.is_some_and(|value| !value.is_empty()),
affinity_lookup_enabled,
);
policy.commit_prefill_selection(&selection_ctx, proposal.kind, &decision.selected);
ctx.metrics
.record_policy_decision(&ctx.config.model.policy.to_string(), reason);
tracing::debug!(
model = %model_str,
policy = ?proposal.kind,
range = %decision.candidate_range_id,
primary = %decision.primary.url,
backup = ?decision.backup.as_ref().map(|worker| worker.url.as_str()),
selected = %decision.selected.url,
reason = ?decision.reason,
load_snapshot_version = decision.load_snapshot_version,
"prefill policy decision",
);
Some(decision.selected)
} else {
tracing::debug!(
model = %model_str,
policy = ?proposal.kind,
range = %candidate_range.id,
selected = %proposal.primary.url,
"prefill policy decision without shared admission",
);
Some(proposal.primary)
}
};
let select_prefill_domains =
|domains: &[CandidateDomain],
affinity_lookup_enabled: bool,
affinity_assignment_enabled: bool| {
domains
.iter()
.find_map(|domain| {
select_prefill_in_domain(
domain,
affinity_lookup_enabled,
affinity_assignment_enabled,
false,
)
})
.or_else(|| {
domains.iter().find_map(|domain| {
select_prefill_in_domain(
domain,
affinity_lookup_enabled,
affinity_assignment_enabled,
true,
)
})
})
};
// Cache-Aware resolves one bounded global candidate set and returns a final winner.
let cache_winner = (ctx.config.model.policy == PolicyKind::CacheAware)
.then(|| {
let snapshot = load_snapshot.as_ref()?;
let global_range = CandidateRange::global(&workers);
let cache_ctx =
SelectionContext::with_routing_key(&model_id, Some(&body), routing_key)
.with_session_id(session_id)
.with_candidate_range_id(global_range.id)
.with_input_tokens(request_input_tokens)
.with_request_tokens(
request_tokens.as_ref().map(|tokens| tokens.ids.as_slice()),
)
.with_external_prefix(external_prefix.as_ref())
.with_load_snapshot(snapshot)
.with_prefill_cache_bucket(&ctx.bucket_selector, prefill_bucket_request);
let PrefillProposal::CacheCandidates(proposal) =
policy.propose_prefill(global_range.workers, &cache_ctx)?
else {
return None;
};
let bounded_candidate_count = proposal.candidates.len();
let cache_decision =
resolve_cache_candidates(&proposal, request_input_tokens, snapshot);
ctx.metrics.record_cache_admission_evaluations(
cache_decision.admission_evaluated_candidates,
);
ctx.metrics.record_cache_admission_rejections(
cache_decision.admission_rejected_candidates,
);
ctx.metrics.record_cache_pressure_guard(
cache_decision.pressure_guard_compared_pairs,
cache_decision.pressure_guard_overrides,
);
ctx.metrics
.record_cache_monitor_decision(cache_decision.prefill_pressure_source);
let Some(decision) = cache_decision.decision else {
selection_failure_reason
.set(PolicySelectionFailureReason::CacheCandidatesExhausted);
return None;
};
let selected_candidate = proposal
.candidates
.iter()
.find(|candidate| candidate.worker.id == decision.selected.id)?;
tracing::debug!(
model = %model_str,
policy = ?ProposalKind::CacheAffinity,
range = %decision.candidate_range_id,
selected = %decision.selected.url,
cache_candidates = bounded_candidate_count,
input_tokens = request_input_tokens,
matched_prefix_tokens = selected_candidate.matched_prefix_tokens,
uncached_tokens = selected_candidate.uncached_tokens,
reason = ?decision.reason,
load_snapshot_version = decision.load_snapshot_version,
prefill_pressure_source = cache_decision.prefill_pressure_source,
"cache candidate winner",
);
ctx.metrics
.record_policy_decision("cache_aware", "cache_candidate");
Some(decision.selected)
})
.flatten();
let global_affinity_probe = use_global_affinity_probe
.then(|| {
let snapshot = load_snapshot.as_ref()?;
let global_range = CandidateRange::global(&workers);
let probe_ctx =
SelectionContext::with_routing_key(&model_id, Some(&body), routing_key)
.with_session_id(session_id)
.with_candidate_range_id(global_range.id)
.with_input_tokens(request_input_tokens)
.with_request_tokens(
request_tokens.as_ref().map(|tokens| tokens.ids.as_slice()),
)
.with_external_prefix(external_prefix.as_ref())
.with_load_snapshot(snapshot)
.without_affinity_assignment();
policy.propose(global_range.workers, &probe_ctx)
})
.flatten();
// A new or stale session may create its first assignment in the target Bucket.
let global_affinity_missed = global_affinity_probe
.as_ref()
.is_some_and(|proposal| !matches!(proposal.kind, ProposalKind::SessionAffinity));
let global_affinity_worker = global_affinity_probe
.and_then(|proposal| {
matches!(proposal.kind, ProposalKind::SessionAffinity).then_some(proposal.primary)
})
.and_then(|primary| {
ctx.bucket_selector.prefill_affinity_domain(
&workers,
&primary,
prefill_bucket_request,
)
})
// Rebuild the backup inside the primary's own Bucket.
.and_then(|domain| select_prefill_in_domain(&domain, true, false, false));
cache_winner
.or_else(|| {
// Materialize normal domains only when Cache-Aware has no winner.
let prefill_domains = ctx
.bucket_selector
.prefill_domains(&workers, prefill_bucket_request);
if ctx.config.model.policy == PolicyKind::CacheAware {
// Cache miss or failure retries ordered domains with ordinary P2.
return select_prefill_domains(&prefill_domains, false, false);
}
global_affinity_worker.or_else(|| match session_affinity_mode {
SessionAffinityMode::GlobalPreserve if global_affinity_missed => {
select_prefill_domains(&prefill_domains, true, true)
}
SessionAffinityMode::GlobalPreserve => {
select_prefill_domains(&prefill_domains, false, false)
}
SessionAffinityMode::Bucket | SessionAffinityMode::GlobalRebind => {
select_prefill_domains(&prefill_domains, true, true)
}
})
})
.ok_or_else(|| {
policy_selection_failed(&ctx, &model_str, selection_failure_reason.get())
})?
};
// PD-mode decoder affinity. When the selected prefill worker is
// part of a PD-disagg deployment, also resolve the matching decode
// peer (same host where possible, falling back to min-load via
// `select_decode_with_affinity`). Both workers receive the SAME
// request body — augmented with the three flat `bootstrap_*`
// fields below — so the SGLang engine can match incoming KV
// transfers via `bootstrap_room`.
// Decode selection starts after Final P.
//
// Plain-mode workers skip the decode resolution entirely (no
// decode peer to find). PD-mode requests that fail to resolve a
// decode peer (`NoDecodeWorkersAvailable`) bubble up as 503 so
// operators can alert on prefill-vs-decode pool imbalance.
let decode_peer: Option<Arc<Worker>> = if worker.mode() == WorkerMode::Prefill {
Some(
resolver
.decode_with_affinity(&model_id, &worker.url)
.map_err(|e| match e {
PdResolveError::NoHealthyWorkers => ApiError::NoHealthyWorkers {
model: model_str.clone(),
},
PdResolveError::NoDecodeWorkersAvailable => {
ApiError::NoDecodeWorkersAvailable {
model: model_str.clone(),
}
}
PdResolveError::NoPrefillWorkersAvailable => {
ApiError::NoPrefillWorkersAvailable {
model: model_str.clone(),
}
}
})?,
)
let decode_workers = resolver.decode_candidates(&model_id).map_err(|e| match e {
PdResolveError::NoHealthyWorkers => ApiError::NoHealthyWorkers {
model: model_str.clone(),
},
PdResolveError::NoDecodeWorkersAvailable => ApiError::NoDecodeWorkersAvailable {
model: model_str.clone(),
},
PdResolveError::NoPrefillWorkersAvailable => ApiError::NoPrefillWorkersAvailable {
model: model_str.clone(),
},
})?;
let request_kv_tokens =
projected_decode_kv_tokens(request_input_tokens, requested_max_output_tokens);
let expected_peak_sequence_tokens = requested_max_output_tokens.map(|_| request_kv_tokens);
let decode_domains = ctx.bucket_selector.decode_domains(
&decode_workers,
BucketRequest {
input_tokens: request_input_tokens,
expected_peak_sequence_tokens,
ttft_slo_ms,
tps_slo,
},
);
let decode_policy = build_decode_policy(ctx.config.model.decode_policy);
let select_decode_in_domain =
|decode_domain: &CandidateDomain, allow_capacity_fallback: bool| {
let snapshot = load_snapshot.as_ref()?;
let decode_ctx = DecodeSelectionContext::new()
.with_load_snapshot(snapshot)
.with_prefill_url(&worker.url);
let decode_proposal = decode_policy.propose(decode_domain, &decode_ctx)?;
let decode_decision = if allow_capacity_fallback {
resolve_decode_with_capacity_fallback(
decode_domain,
&decode_proposal,
request_kv_tokens,
snapshot,
)
} else {
resolve_decode(decode_domain, &decode_proposal, request_kv_tokens, snapshot)
}?;
tracing::debug!(
model = %model_str,
policy = ?ctx.config.model.decode_policy,
range = %decode_decision.candidate_range_id,
primary = %decode_decision.primary.url,
backup = ?decode_decision.backup.as_ref().map(|worker| worker.url.as_str()),
selected = %decode_decision.selected.url,
reason = ?decode_decision.reason,
load_snapshot_version = decode_decision.load_snapshot_version,
"decode policy decision",
);
Some(decode_decision.selected)
};
decode_domains
.iter()
.find_map(|domain| select_decode_in_domain(domain, false))
.or_else(|| {
decode_domains
.iter()
.find_map(|domain| select_decode_in_domain(domain, true))
})
.ok_or_else(|| ApiError::NoDecodeWorkersAvailable {
model: model_str.clone(),
})
.map(Some)?
} else {
None
};
@@ -373,11 +718,8 @@ pub async fn chat_completions(
// ends, the client disconnects, or the handler returns an error. In
// PD mode the pair moves into the spawned prefill task so prefill
// load is tracked for the full duration of the KV transfer; in plain
// mode the pair stays in this handler. Decode-load contribution is
// 0 here: the active-load registry's decode axis is reserved for a
// future decode-side scheduler — current decode selection is
// host-affinity only.
let guard = if needs_load_snapshot {
// mode the pair stays in this handler. Decode load is tracked on Final D.
let guard = if needs_dispatch_timestamps {
worker.timestamped_load_guard()
} else {
worker.load_guard()
@@ -551,13 +893,14 @@ pub async fn chat_completions(
// Synchronously await the decode worker. Its response is what
// the client sees. The decode side gets its own LoadGuard so
// per-worker `active_requests` reflects decode-pool load. Decode
// selection reads that atomic counter directly, so it does not need
// the prefill policy's timestamp registry.
// per-worker `active_requests` reflects load on Final D.
let decode_guard = decode_worker.load_guard();
let decode_active_guard =
ctx.active_load
.register(decode_worker.id.clone(), decode_worker.url.clone(), 0, 1);
if streaming {
let stream_guards: Box<dyn Send + 'static> =
Box::new((decode_guard, make_duration_guard()));
Box::new((decode_guard, decode_active_guard, make_duration_guard()));
let fetch = ctx.proxy.forward_streaming_to(
&decode_worker.url,
&decode_worker.breaker,
@@ -573,7 +916,7 @@ pub async fn chat_completions(
_ = stale_token.cancelled() => Err(ApiError::StaleRequestExpired { model: model_str }),
}
} else {
let _decode_hold = decode_guard;
let _decode_hold = (decode_guard, decode_active_guard);
let fetch = ctx.proxy.forward_json_to(
&decode_worker.url,
&decode_worker.breaker,
@@ -706,7 +1049,7 @@ pub async fn chat_completions(
);
// Mirror the upstream `x-sgl-decode-url` hint onto the response so
// external tests / sidecars can observe PD decode affinity without
// external tests / sidecars can observe the final PD Decode selection without
// sniffing the proxy hop. The request-side header was set above for
// the prefill worker; copying it here makes the affinity observable
// end-to-end. Plain-mode requests skip this (no decode peer was
@@ -770,17 +1113,65 @@ fn resolve_prefix_query(
}
}
fn parse_optional_positive_u64_header(
headers: &HeaderMap,
name: &HeaderName,
label: &str,
) -> Result<Option<u64>, ApiError> {
let Some(value) = headers.get(name) else {
return Ok(None);
};
let raw = value
.to_str()
.map_err(|_| ApiError::BadRequest(format!("{label} header must be ASCII")))?;
let parsed = raw
.parse::<u64>()
.map_err(|_| ApiError::BadRequest(format!("{label} header must be a positive integer")))?;
if parsed == 0 {
return Err(ApiError::BadRequest(format!(
"{label} header must be a positive integer"
)));
}
Ok(Some(parsed))
}
fn parse_optional_positive_f64_header(
headers: &HeaderMap,
name: &HeaderName,
label: &str,
) -> Result<Option<f64>, ApiError> {
let Some(value) = headers.get(name) else {
return Ok(None);
};
let raw = value
.to_str()
.map_err(|_| ApiError::BadRequest(format!("{label} header must be ASCII")))?;
let parsed = raw
.parse::<f64>()
.map_err(|_| ApiError::BadRequest(format!("{label} header must be a positive number")))?;
if !parsed.is_finite() || parsed <= 0.0 {
return Err(ApiError::BadRequest(format!(
"{label} header must be a finite positive number"
)));
}
Ok(Some(parsed))
}
fn should_tokenize_request(
has_chat_encoder: bool,
policy_needs_request_tokens: bool,
bucket_enabled: bool,
) -> bool {
has_chat_encoder || policy_needs_request_tokens || bucket_enabled
}
/// Estimate prefill-token count from the raw request body for use as
/// the active-load `prefill_load` counter. Returns 1 at minimum so
/// a registered request always shows up as "load > 0" — under-counting
/// to zero would hide the request from the cache-aware policy's
/// load-imbalance fast-path.
///
/// This is a coarse approximation: we count the body length in bytes
/// and divide by [`CHARS_PER_TOKEN_ESTIMATE`]. A future improvement is
/// to thread the tokenizer's actual token count through (the
/// cache-aware-zmq policy already tokenizes the prompt for tree
/// matching — that count could be reused here).
/// Exact ingress tokens are preferred when available.
fn estimate_prefill_tokens(body: &Bytes) -> usize {
(body.len() / CHARS_PER_TOKEN_ESTIMATE).max(1)
}
@@ -857,7 +1248,7 @@ fn build_outgoing_body(
_ => {
return Err(ApiError::BadRequest(
"invalid request: body must be a JSON object".to_string(),
))
));
}
};
if let Some(ids) = input_ids {
@@ -1093,6 +1484,93 @@ mod tests {
));
}
#[test]
fn bucket_routing_requests_tokens_even_for_a_non_token_policy() {
assert!(should_tokenize_request(false, false, true));
assert!(!should_tokenize_request(false, false, false));
}
#[test]
fn session_reason_distinguishes_hit_assignment_and_keyless_fallback() {
assert_eq!(
prefill_policy_reason(
PolicyKind::SessionAware,
ProposalKind::SessionAffinity,
DecisionReason::Primary,
true,
true,
),
"session_primary"
);
assert_eq!(
prefill_policy_reason(
PolicyKind::SessionAware,
ProposalKind::PowerOfTwo,
DecisionReason::Primary,
true,
true,
),
"assigned"
);
assert_eq!(
prefill_policy_reason(
PolicyKind::SessionAware,
ProposalKind::PowerOfTwo,
DecisionReason::Primary,
false,
true,
),
"no_session"
);
}
#[test]
fn session_reason_preserves_admission_and_pressure_escapes() {
assert_eq!(
prefill_policy_reason(
PolicyKind::SessionAware,
ProposalKind::SessionAffinity,
DecisionReason::BackupPrimaryAdmission,
true,
true,
),
"session_admission_backup"
);
assert_eq!(
prefill_policy_reason(
PolicyKind::SessionAware,
ProposalKind::SessionAffinity,
DecisionReason::BackupPressureGuard,
true,
true,
),
"session_pressure_backup"
);
}
#[test]
fn cache_no_winner_p2_is_distinct_from_cache_candidate() {
assert_eq!(
prefill_policy_reason(
PolicyKind::CacheAware,
ProposalKind::PowerOfTwo,
DecisionReason::Primary,
false,
false,
),
"no_cache_candidate"
);
assert_eq!(
prefill_policy_reason(
PolicyKind::CacheAware,
ProposalKind::CacheAffinity,
DecisionReason::Primary,
false,
true,
),
"cache_candidate"
);
}
/// `generate_room_id` MUST return values in `[0, i64::MAX]`. The
/// SGLang prefill stores `bootstrap_room` as `torch.int64`; a u64
/// with the top bit set would wrap negative on the engine side.
@@ -1361,6 +1839,33 @@ mod tests {
assert_eq!(p.model.as_deref(), Some("tiny"));
}
#[test]
fn parse_probe_accepts_modern_openai_completion_budget() {
let body =
Bytes::from_static(br#"{"model":"tiny","messages":[],"max_completion_tokens":256}"#);
assert_eq!(
parse_probe(&body).unwrap().requested_max_output_tokens(),
Some(256)
);
}
#[test]
fn modern_completion_budget_takes_precedence_when_both_fields_are_present() {
let body =
Bytes::from_static(br#"{"model":"tiny","max_tokens":128,"max_completion_tokens":256}"#);
assert_eq!(
parse_probe(&body).unwrap().requested_max_output_tokens(),
Some(256)
);
}
#[test]
fn decode_kv_projection_includes_the_explicit_output_budget() {
assert_eq!(projected_decode_kv_tokens(1_024, Some(512)), 1_536);
assert_eq!(projected_decode_kv_tokens(1_024, None), 1_024);
assert_eq!(projected_decode_kv_tokens(u64::MAX - 1, Some(8)), u64::MAX);
}
#[test]
fn parse_probe_rejects_non_object_shapes() {
// Pin the contract: degenerate JSON (valid JSON but wrong shape)
@@ -95,7 +95,6 @@ mod tests {
let body = std::str::from_utf8(&body).unwrap();
// Every metric family should at least carry its HELP/TYPE lines.
assert!(body.contains("# TYPE sgl_router_requests_total counter"));
assert!(body.contains("# TYPE sgl_router_overlap_blocks histogram"));
assert!(body.contains("# TYPE sgl_router_active_load gauge"));
}
@@ -51,6 +51,8 @@ mod tests {
id: "qwen3".into(),
tokenizer_path: "x".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -118,6 +118,8 @@ mod tests {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -6,9 +6,8 @@
//! The engine caches KV blocks keyed on tokens it produces *after* applying the
//! model's chat template (BOS + role/special markers + content). The router's
//! cache-aware selection must hash the same token sequence, so it renders the
//! same template before tokenizing otherwise its query hashes never match the
//! engine's stored blocks and cache-aware routing silently degrades to min-load
//! (`sgl_router_overlap_blocks_sum` stuck at 0).
//! same template before tokenizing; otherwise its query hashes never match the
//! engine's stored blocks.
//!
//! The template and its special-token strings come from the model's
//! `tokenizer_config.json` — the HuggingFace built-in template, which is what
@@ -232,6 +232,8 @@ mod tests {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -51,8 +51,7 @@ pub async fn run(rx: mpsc::Receiver<DiscoveryEvent>, registry: Arc<WorkerRegistr
/// on every worker add / remove, and an optional active-load registry
/// that is asked to forget per-worker counters on `Removed`.
///
/// When `kv_index` is `None` the cache-aware-zmq path is disabled
/// (selection falls through to the non-cache-aware policies); when
/// When `kv_index` is `None`, KV-event and load-subscriber state is disabled; when
/// `active_load` is `None` the active-load bookkeeping is not pruned
/// on worker removal (leaks one `WorkerCounters` slot per departed
/// worker — fine for tests, but production passes `Some(...)`); when
@@ -475,6 +474,8 @@ mod tests {
id: id.into(),
tokenizer_path: "/tmp/x".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: Some(RawCbConfig {
threshold: NonZeroU32::new(threshold).unwrap(),
cool_down_secs,
@@ -67,9 +67,8 @@ impl SlotRegistry {
/// Count of currently-claimed slots acquired at or after `since`. Used to
/// bound how many of this worker's in-flight requests are dispatches the
/// engine hasn't reported back on yet (see
/// `crate::policies::cache_aware_zmq::WorkerLoads::load_of`), rather than
/// adding the full in-flight count — which would also include long-held
/// engine hasn't reported back on yet, rather than adding the full
/// in-flight count — which would also include long-held
/// slots from slow-draining streaming responses (see
/// `crate::proxy::Proxy::forward_streaming_to`'s `stream_guards` doc)
/// that the engine's own last report likely already accounts for.
@@ -130,6 +130,8 @@ async fn static_urls_pd_role_resolved_end_to_end() {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: sgl_router::config::PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -0,0 +1,356 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Static buckets define candidate domains and fallback order. Worker scoring,
//! admission, and guards remain the responsibility of the P/D policies.
use sgl_router::config::{BucketConfig, BucketSpec, BucketStage, SloBucketPolicy};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::buckets::{BucketRequest, BucketSelector};
use sgl_router::policies::CacheCandidate;
use sgl_router::workers::Worker;
use std::sync::Arc;
fn worker(id: &str, mode: WorkerMode) -> Arc<Worker> {
Arc::new(Worker::new(WorkerSpec {
id: WorkerId(id.into()),
url: format!("http://{id}:30000"),
mode,
model_ids: vec![ModelId("m".into())],
bootstrap_port: None,
}))
}
fn bucket(id: &str, stage: BucketStage, rank: u32, worker_ids: &[&str]) -> BucketSpec {
BucketSpec {
id: id.into(),
stage,
rank,
worker_ids: worker_ids.iter().map(|id| (*id).into()).collect(),
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: None,
ttft_p95_at_capacity_ms: None,
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
}
}
#[test]
fn prefill_slo_first_tries_eligible_buckets_by_rank_before_degrading() {
let fast = worker("fast", WorkerMode::Prefill);
let cheap = worker("cheap", WorkerMode::Prefill);
let mut cheap_bucket = bucket("cheap", BucketStage::Prefill, 10, &["cheap"]);
cheap_bucket.ttft_p95_at_capacity_ms = Some(300);
let mut fast_bucket = bucket("fast", BucketStage::Prefill, 20, &["fast"]);
fast_bucket.ttft_p95_at_capacity_ms = Some(80);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![cheap_bucket, fast_bucket],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let domains = selector.prefill_domains(
&[cheap, fast],
BucketRequest {
input_tokens: 256,
expected_peak_sequence_tokens: None,
ttft_slo_ms: Some(100),
tps_slo: None,
},
);
assert_eq!(
domains
.iter()
.map(|domain| domain.id.as_str())
.collect::<Vec<_>>(),
["fast", "cheap"],
"eligible buckets come first; non-eligible buckets are the explicit SLO-degraded fallback"
);
}
#[test]
fn prefill_best_effort_tries_non_slo_bucket_before_reserved_slo_capacity() {
let fast = worker("fast", WorkerMode::Prefill);
let cheap = worker("cheap", WorkerMode::Prefill);
let mut fast_bucket = bucket("fast", BucketStage::Prefill, 10, &["fast"]);
fast_bucket.ttft_p95_at_capacity_ms = Some(80);
let mut cheap_bucket = bucket("cheap", BucketStage::Prefill, 20, &["cheap"]);
cheap_bucket.ttft_p95_at_capacity_ms = Some(300);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![fast_bucket, cheap_bucket],
ttft_slo_policy: SloBucketPolicy::BestEffort,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let domains = selector.prefill_domains(
&[fast, cheap],
BucketRequest {
input_tokens: 256,
expected_peak_sequence_tokens: None,
ttft_slo_ms: Some(100),
tps_slo: None,
},
);
assert_eq!(
domains
.iter()
.map(|domain| domain.id.as_str())
.collect::<Vec<_>>(),
["cheap", "fast"],
"best-effort tries non-SLO capacity first and retains the SLO tier as fallback"
);
}
#[test]
fn cache_candidate_uses_uncached_work_range_but_full_context_and_own_ttft_profile() {
let short = worker("short", WorkerMode::Prefill);
let long = worker("long", WorkerMode::Prefill);
let mut short_bucket = bucket("p-short", BucketStage::Prefill, 10, &["short"]);
short_bucket.max_extend_tokens = Some(64);
short_bucket.max_context_tokens = Some(4_096);
short_bucket.ttft_p95_at_capacity_ms = Some(80);
let mut long_bucket = bucket("p-long", BucketStage::Prefill, 20, &["long"]);
long_bucket.min_extend_tokens = Some(65);
long_bucket.max_context_tokens = Some(4_096);
long_bucket.ttft_p95_at_capacity_ms = Some(300);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![short_bucket, long_bucket],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let workers = vec![Arc::clone(&short), Arc::clone(&long)];
let request = BucketRequest {
input_tokens: 256,
expected_peak_sequence_tokens: None,
ttft_slo_ms: Some(100),
tps_slo: None,
};
assert_eq!(
selector
.prefill_domains(&workers, request)
.iter()
.map(|domain| domain.id.as_str())
.collect::<Vec<_>>(),
["p-long"],
"no-hit target selection uses E=L for extend-work compatibility"
);
let short_hit = CacheCandidate {
worker: Arc::clone(&short),
matched_prefix_tokens: 224,
uncached_tokens: 32,
candidate_range_id: "global".into(),
max_pending_prefill_tokens: None,
};
let bound = selector
.bind_prefill_cache_candidate(short_hit, request)
.expect("E=32 fits short work range and the full L=256 fits max context");
assert_eq!(bound.candidate_range_id, "p-short");
let long_hit = CacheCandidate {
worker: Arc::clone(&long),
matched_prefix_tokens: 0,
uncached_tokens: 256,
candidate_range_id: "global".into(),
max_pending_prefill_tokens: None,
};
assert!(
selector
.bind_prefill_cache_candidate(long_hit, request)
.is_none(),
"a cache candidate whose own Hard TTFT profile misses the request SLO is rejected"
);
}
#[test]
fn cache_candidate_without_bucket_configuration_keeps_global_metadata() {
let p = worker("p", WorkerMode::Prefill);
let selector = BucketSelector::new(None);
let candidate = CacheCandidate {
worker: p,
matched_prefix_tokens: 64,
uncached_tokens: 64,
candidate_range_id: "probe".into(),
max_pending_prefill_tokens: Some(1),
};
let bound = selector
.bind_prefill_cache_candidate(
candidate,
BucketRequest {
input_tokens: 128,
expected_peak_sequence_tokens: None,
ttft_slo_ms: None,
tps_slo: None,
},
)
.expect("Step 1 always has a catch-all domain");
assert_eq!(bound.candidate_range_id, "global");
assert_eq!(bound.max_pending_prefill_tokens, None);
}
#[test]
fn decode_bucket_uses_peak_sequence_length_then_tps_profile_and_rank() {
let short = worker("short", WorkerMode::Decode);
let long = worker("long", WorkerMode::Decode);
let mut short_bucket = bucket("short", BucketStage::Decode, 10, &["short"]);
short_bucket.max_sequence_tokens = Some(1_024);
short_bucket.tps_p05_at_capacity = Some(80.0);
let mut long_bucket = bucket("long", BucketStage::Decode, 20, &["long"]);
long_bucket.max_sequence_tokens = Some(8_192);
long_bucket.tps_p05_at_capacity = Some(40.0);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![short_bucket, long_bucket],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::SloFirst,
}));
let domains = selector.decode_domains(
&[short, long],
BucketRequest {
input_tokens: 256,
expected_peak_sequence_tokens: Some(900),
ttft_slo_ms: None,
tps_slo: Some(60.0),
},
);
assert_eq!(domains.len(), 2);
assert_eq!(domains[0].id, "short");
assert_eq!(domains[1].id, "long");
}
#[test]
fn missing_bucket_configuration_keeps_the_global_domain() {
let p = worker("p", WorkerMode::Prefill);
let d = worker("d", WorkerMode::Decode);
let selector = BucketSelector::new(None);
let facts = BucketRequest {
input_tokens: 128,
expected_peak_sequence_tokens: Some(512),
ttft_slo_ms: Some(100),
tps_slo: Some(20.0),
};
let prefill = selector.prefill_domains(&[p], facts);
let decode = selector.decode_domains(&[d], facts);
assert_eq!(prefill.len(), 1);
assert_eq!(prefill[0].id, "global");
assert_eq!(decode.len(), 1);
assert_eq!(decode[0].id, "global");
}
#[test]
fn prefill_only_bucket_configuration_keeps_the_global_decode_domain() {
let p = worker("p", WorkerMode::Prefill);
let d = worker("d", WorkerMode::Decode);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![bucket("p", BucketStage::Prefill, 10, &["p"])],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let facts = BucketRequest {
input_tokens: 128,
expected_peak_sequence_tokens: Some(512),
ttft_slo_ms: None,
tps_slo: None,
};
assert_eq!(selector.prefill_domains(&[p], facts)[0].id, "p");
let decode = selector.decode_domains(&[d], facts);
assert_eq!(decode.len(), 1);
assert_eq!(decode[0].id, "global");
}
#[test]
fn decode_catch_all_still_rejects_input_beyond_runtime_context() {
let d = worker("d", WorkerMode::Decode);
let mut catch_all = bucket("d-catch-all", BucketStage::Decode, 10, &["d"]);
catch_all.max_context_tokens = Some(1_024);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![catch_all],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let domains = selector.decode_domains(
&[d],
BucketRequest {
input_tokens: 2_048,
expected_peak_sequence_tokens: None,
ttft_slo_ms: None,
tps_slo: None,
},
);
assert!(
domains.is_empty(),
"an unknown output budget does not erase the known input context requirement"
);
}
#[test]
fn membership_index_preserves_exact_matching_and_fleet_order() {
let workers: Vec<_> = (0..10)
.map(|index| worker(&format!("w{index}"), WorkerMode::Prefill))
.collect();
let scan = bucket("scan", BucketStage::Prefill, 10, &["w3", "W3", "w1", "w1"]);
let set = bucket(
"set",
BucketStage::Prefill,
20,
&[
"w9", "w3", "w1", "w1", "W3", " w2", "absent-0", "absent-1", "absent-2",
],
);
let selector = BucketSelector::new(Some(BucketConfig {
buckets: vec![scan, set],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
}));
let request = BucketRequest {
input_tokens: 128,
expected_peak_sequence_tokens: None,
ttft_slo_ms: None,
tps_slo: None,
};
let domains = selector.prefill_domains(&workers, request);
let ids = |index: usize| {
domains[index]
.workers
.iter()
.map(|worker| worker.id.0.as_str())
.collect::<Vec<_>>()
};
assert_eq!(ids(0), ["w1", "w3"]);
assert_eq!(ids(1), ["w1", "w3", "w9"]);
let candidate = CacheCandidate {
worker: Arc::clone(&workers[9]),
matched_prefix_tokens: 0,
uncached_tokens: 128,
candidate_range_id: "global".into(),
max_pending_prefill_tokens: None,
};
assert_eq!(
selector
.bind_prefill_cache_candidate(candidate, request)
.expect("w9 belongs to the hash-indexed bucket")
.candidate_range_id,
"set"
);
assert_eq!(
selector
.prefill_affinity_domain(&workers, &workers[9], request)
.expect("w9 has a bucket affinity")
.id,
"set"
);
}
@@ -1,190 +0,0 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! E2E test for the cache-aware-zmq policy.
//!
//! Drives a real ZMQ PUB socket → `KvEventIndex` subscriber pipeline →
//! `HashTree` → `CacheAwareZmqPolicy::select`. Verifies that an event
//! published by one worker's PUB causes subsequent selection to route
//! to that worker (cache-aware affinity).
//!
//! API constraint: the subscriber registry builds endpoints as
//! `tcp://{host}:{port_base + dp_rank}` where `port_base` is in the
//! per-worker `EventConfig`. Both mock workers below share
//! `127.0.0.1` as host, so both subscribe to the same PUB socket and
//! both end up indexed in the tree. The tiebreak (lowest active_load)
//! picks the worker we want; same shape as the SMG version of this
//! test.
use std::sync::Arc;
use std::time::Duration;
use zeromq::SocketSend;
use sgl_router::config::CacheAwareConfig;
use sgl_router::config::{ActiveLoadConfig, ProxyConfig};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::cache_aware_zmq::CacheAwareZmqPolicy;
use sgl_router::policies::engine_load::EngineLoadTable;
use sgl_router::policies::kv_events::{compute_block_hashes, discovery::EventConfig, KvEventIndex};
use sgl_router::policies::{Policy, SelectionContext};
use sgl_router::tokenizer::TokenizerRegistry;
use sgl_router::workers::Worker;
use super::zmq_helpers::{
build_multipart, encode_block_stored_event, encode_event_batch, make_pub_bound,
};
fn build_worker(url: &str, model: &str) -> Arc<Worker> {
Arc::new(Worker::new(WorkerSpec {
id: WorkerId(url.into()),
url: url.into(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId(model.into())],
bootstrap_port: None,
}))
}
/// E2E: real PUB socket publishes a `BlockStored` for worker A's
/// hash chain. The `CacheAwareZmqPolicy`'s shared `KvEventIndex`
/// receives it, applies it to the tree, and the next `select` call
/// picks worker A.
///
/// Both workers share `127.0.0.1` as host so both subscribers connect
/// to the same PUB and both get indexed under their KvWorkerIds — the
/// same shape as the SMG e2e test. We tie-break on min-load: worker B
/// is bumped above worker A so the matched-worker pick prefers A.
#[tokio::test]
async fn zmq_indexer_routes_to_publishing_worker_e2e() {
let model_id = ModelId("tiny".into());
// 1. Tokenizer registry — use the in-tree tiny fixture.
let cfg = sgl_router::config::Config {
server: sgl_router::config::ServerConfig {
host: "0".into(),
port: 0,
},
observability: Default::default(),
model: sgl_router::config::ModelConfig {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: sgl_router::config::PolicyKind::CacheAwareZmq,
circuit_breaker: None,
cache_aware: None,
sticky: None,
affinity: None,
fused: None,
eligibility: None,
},
discovery: sgl_router::config::DiscoveryBackend::StaticUrls(
sgl_router::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()],
},
),
proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(),
};
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
// 2. Bind a real PUB socket on an OS-assigned port.
let (mut pub_a, port) = make_pub_bound().await;
// 3. Compute the hash chain for the routing prompt.
let text = "hello world hello world hello world";
let tok = tokenizers.get("tiny").unwrap();
let token_ids = sgl_router::tokenizer::adapter::encode(&tok, text).unwrap();
let block_size = 4u32;
let hashes = compute_block_hashes(&token_ids, block_size as usize);
assert!(!hashes.is_empty(), "tiny tokenizer must yield ≥1 block");
// 4. Build the KvEventIndex + policy. The policy holds an
// Arc<HashTree> that the index also owns; events the index
// receives mutate the same tree the policy reads.
let kv_index = KvEventIndex::new();
// Mirror what `KvEventIndex::add_worker` would do in production: seed
// the oracle with the worker-reported page_size before any cache
// lookup happens. The integration path calls `add_worker` further
// down, but here we want the policy to know `block_size` immediately.
let block_size_oracle = kv_index.block_size_oracle();
block_size_oracle.try_set(block_size).unwrap();
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 0.0,
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
kv_indexer_endpoint: None,
},
kv_index.tree(),
Arc::clone(&tokenizers),
block_size_oracle,
EngineLoadTable::new(),
);
// 5. Register two workers. They share `127.0.0.1` so both
// subscribers connect to the same PUB; preresolved EventConfig
// points at the bound port.
let url_a = "http://127.0.0.1:30000";
let url_b = "http://127.0.0.1:30001";
let preresolved = EventConfig {
host: "127.0.0.1".to_string(),
port_base: port,
topic: String::new(),
block_size,
dp_size: 1,
load_port_base: None,
load_topic: None,
is_bigram: false,
};
kv_index.add_worker(url_a, Some(preresolved.clone())).await;
kv_index.add_worker(url_b, Some(preresolved)).await;
// SUB sockets take a moment to handshake. The polling loop below
// soaks up any extra latency; this is just a publish-before-SUB
// guard.
tokio::time::sleep(Duration::from_millis(150)).await;
// 6. Publish a BlockStored event for the routing prompt's chain.
let event_bytes = encode_block_stored_event(&hashes, None, &token_ids, block_size);
let payload = encode_event_batch(0.0, vec![event_bytes], Some(0));
pub_a
.send(build_multipart(1, payload))
.await
.expect("send block-stored event");
// 7. Bump worker B's load so the tie-break picks A among matched
// workers. The bump stays below balance_abs_threshold so the
// imbalance fast-path does not skip cache-aware selection.
// Bind the guards to a Vec held for the rest of the test scope
// so the counter stays > 0 through the polling loop.
let w_a = build_worker(url_a, "tiny");
let w_b = build_worker(url_b, "tiny");
let _b_load: Vec<_> = (0..3).map(|_| w_b.load_guard()).collect();
let workers = vec![Arc::clone(&w_a), Arc::clone(&w_b)];
// 8. Drive select until the event has been applied. The pipeline is
// asynchronous (publish → SUB recv → mpsc → pump → tree); a
// polling loop is less flaky than a fixed sleep.
let body = serde_json::to_vec(&serde_json::json!({"prompt": text})).unwrap();
let ctx = SelectionContext::new(&model_id, Some(&body));
let start = std::time::Instant::now();
let mut chose_a = false;
while start.elapsed() < Duration::from_secs(3) {
if let Some(w) = policy.select(&workers, &ctx) {
if w.url == url_a {
chose_a = true;
break;
}
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
assert!(
chose_a,
"policy did not route to publishing worker A within timeout",
);
// 9. Shutdown cleanly.
let r = tokio::time::timeout(Duration::from_secs(2), kv_index.shutdown()).await;
assert!(r.is_ok(), "kv_index shutdown should not hang");
}
@@ -0,0 +1,52 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use std::collections::HashMap;
use std::sync::Arc;
use sgl_kv_indexer::PrefixOutcome;
use sgl_router::policies::kv_events::{
compute_block_hashes, BlockSizeOracle, HashTree, KvWorkerId,
};
use sgl_router::policies::prefix_provider::RadixTreePrefixProvider;
#[test]
fn radix_tree_reports_contiguous_prefix_depth_per_worker() {
let tokens = [11_u32, 12, 13, 14];
let hashes = compute_block_hashes(&tokens, 1);
let tree = Arc::new(HashTree::new());
let oracle = BlockSizeOracle::new();
oracle.try_set(1).unwrap();
tree.insert(&KvWorkerId::new("http://deep".into(), 0), None, &hashes);
tree.insert(
&KvWorkerId::new("http://deep".into(), 1),
None,
&hashes[..3],
);
tree.insert(
&KvWorkerId::new("http://shallow".into(), 0),
None,
&hashes[..2],
);
let signal = RadixTreePrefixProvider::new(tree, oracle)
.match_request_tokens(&tokens)
.expect("established local tree must produce a prefix signal");
let PrefixOutcome::Matched {
matches,
best_prefix_blocks,
} = signal.outcome
else {
panic!("local radix-tree hit must be normalized as a match");
};
let depth_by_url: HashMap<_, _> = matches
.into_iter()
.map(|entry| (entry.address, entry.matched_prefix_blocks))
.collect();
assert_eq!(signal.query_blocks, 4);
assert_eq!(best_prefix_blocks, 4);
assert_eq!(depth_by_url.get("http://deep"), Some(&4));
assert_eq!(depth_by_url.get("http://shallow"), Some(&2));
}
@@ -0,0 +1,153 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Observable contract for decode policies.
//!
//! Decode guards require complete, fresh native monitor samples. Short frames
//! fall back to local load and must not appear as monitor-backed decisions.
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::admission::{resolve_decode, CandidateDomain, DecisionReason};
use sgl_router::policies::decode::{
resolve_decode_with_capacity_fallback, DecodePolicy, DecodePowerOfTwoPolicy,
DecodeSelectionContext, LegacyHostAffinityDecodePolicy,
};
use sgl_router::policies::engine_load::{EngineLoadSnapshot, NativeCacheWorkerLoad};
use sgl_router::policies::SelectionProposal;
use sgl_router::workers::Worker;
use std::collections::HashMap;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::Instant;
fn worker(id: &str) -> Arc<Worker> {
Arc::new(Worker::new(WorkerSpec {
id: WorkerId(id.into()),
url: format!("http://{id}:30000"),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("m".into())],
bootstrap_port: None,
}))
}
fn snapshot(entries: &[(&Arc<Worker>, u64, u64, u64, u64)]) -> EngineLoadSnapshot {
EngineLoadSnapshot::from_native_cache_workers(
7,
entries
.iter()
.map(|(worker, running, waiting, used, capacity)| {
(
worker.url.clone(),
NativeCacheWorkerLoad {
num_running_reqs: *running,
num_waiting_reqs: *waiting,
num_waiting_uncached_tokens: *waiting,
num_used_tokens: *used,
num_total_tokens: *used,
max_total_num_tokens: *capacity,
max_running_requests: 64,
prefill_throughput_tokens_per_s: None,
estimated_prefill_queue_ms: None,
captured_at: Instant::now(),
},
)
})
.collect::<HashMap<_, _>>(),
)
}
#[test]
fn decode_p2_proposes_a_distinct_lower_pressure_primary_and_backup() {
let busy = worker("busy");
let idle = worker("idle");
busy.active_requests.store(8, Ordering::Relaxed);
idle.active_requests.store(1, Ordering::Relaxed);
let domain = CandidateDomain::global_decode(&[Arc::clone(&busy), Arc::clone(&idle)]);
let ctx = DecodeSelectionContext::new();
let proposal = DecodePowerOfTwoPolicy::new()
.propose(&domain, &ctx)
.expect("two decode candidates must produce a proposal");
assert_eq!(proposal.primary.id, idle.id);
assert_eq!(
proposal.backup.expect("P2 keeps the other sample").id,
busy.id
);
}
#[test]
fn legacy_host_affinity_remains_an_explicit_single_primary_compatibility_policy() {
let same_host = worker("host-a");
let other_host = worker("host-b");
let domain = CandidateDomain::global_decode(&[Arc::clone(&same_host), other_host]);
let ctx = DecodeSelectionContext::new().with_prefill_url("http://host-a:9999");
let proposal = LegacyHostAffinityDecodePolicy
.propose(&domain, &ctx)
.expect("legacy policy selects one compatible decode worker");
assert_eq!(proposal.primary.id, same_host.id);
assert!(
proposal.backup.is_none(),
"legacy semantics do not invent a backup"
);
}
#[test]
fn decode_admission_uses_backup_before_scanning_domain() {
let primary = worker("primary");
let backup = worker("backup");
let fallback = worker("fallback");
let domain = CandidateDomain::global_decode(&[
Arc::clone(&primary),
Arc::clone(&backup),
Arc::clone(&fallback),
]);
let loads = snapshot(&[
(&primary, 4, 0, 950, 1_000),
(&backup, 0, 0, 0, 1_000),
(&fallback, 0, 0, 0, 1_000),
]);
let proposal = SelectionProposal::with_backup(Arc::clone(&primary), Arc::clone(&backup));
let decision =
resolve_decode(&domain, &proposal, 64, &loads).expect("admitted backup must be selected");
assert_eq!(decision.selected.id, backup.id);
assert_eq!(decision.reason, DecisionReason::BackupPrimaryAdmission);
}
#[test]
fn decode_guard_can_escape_a_primary_to_lower_dynamic_pressure_backup() {
let primary = worker("primary");
let backup = worker("backup");
let domain = CandidateDomain::global_decode(&[Arc::clone(&primary), Arc::clone(&backup)]);
let loads = snapshot(&[(&primary, 3, 2, 900, 2_000), (&backup, 1, 0, 100, 2_000)]);
let proposal = SelectionProposal::with_backup(Arc::clone(&primary), Arc::clone(&backup));
let decision =
resolve_decode(&domain, &proposal, 64, &loads).expect("both candidates are admitted");
assert_eq!(decision.selected.id, backup.id);
assert_eq!(decision.reason, DecisionReason::BackupPressureGuard);
}
#[test]
fn decode_all_capacity_rejected_falls_back_to_power_of_two_within_domain() {
let primary = worker("primary");
let backup = worker("backup");
let workers = vec![Arc::clone(&primary), Arc::clone(&backup)];
let domain = CandidateDomain::global_decode(&workers);
let loads = snapshot(&[
(&primary, 0, 0, 1_000, 1_000),
(&backup, 0, 10, 1_000, 1_000),
]);
let proposal = SelectionProposal::with_backup(Arc::clone(&primary), Arc::clone(&backup));
let decision = resolve_decode_with_capacity_fallback(&domain, &proposal, 64, &loads)
.expect("capacity exhaustion must degrade within the decode domain");
assert_eq!(decision.selected.id, primary.id);
assert_eq!(decision.reason, DecisionReason::CapacityFallbackPowerOfTwo);
}
@@ -88,7 +88,7 @@ async fn two_independent_subscribers_converge_to_same_tree_state() {
&& mb.workers.contains(&key);
if converged {
// Both trees agree on count AND on the worker that holds the
// prefix. This is what the cache-aware-zmq policy reads to
// prefix. This is what the Radix Tree provider reads to
// pick a worker; both routers picking the same key here
// means they would route the same prompt to the same worker.
assert_eq!(
@@ -3,7 +3,9 @@
mod zmq_helpers;
mod cache_aware_zmq;
mod bucket_domains;
mod cache_prefix_provider;
mod decode;
mod fused_score;
mod kv_events_hash_parity;
mod kv_events_tree_concurrent;
@@ -1,325 +0,0 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
# SPDX-License-Identifier: Apache-2.0
"""Content-based routing test for both cache-aware-zmq index backends.
Two SGLang workers publish KV events to two routers at once: one runs the
local ``KvEventIndex`` (SUB straight to the workers) and one runs against an
external KV Indexer fed by a ``kv-indexer-bridge`` per worker. Every
subscriber attaches before the single warmup, so one pair of disjoint
prefixes exercises both index backends without a second model load.
Assert on content, not on convergence: a broken event path degrades
``cache_aware_zmq`` to content-blind min-load, which routes both prefixes to
one worker and so fails at least one assertion below.
"""
from __future__ import annotations
import os
import re
import socket
import subprocess
import time
from contextlib import contextmanager
from pathlib import Path
import httpx
import pytest
from infra.gateway import Gateway
from infra.model_pool import spawn_worker
from infra.model_specs import get_model_spec
# Disjoint prefixes — share no common content. Under the chat template both
# render with the same leading role header (``<|im_start|>user`` ...; Qwen3 has
# no BOS token), so the first block(s) may hash identically; the disjoint
# content then diverges
# well within the matched region, making each worker's HashTree contribution
# uniquely identifying.
#
# Length matters: each prefix must span ≥2 SGLang blocks at the default
# block_size of 64 tokens so the worker actually emits BlockStored
# events. Below that, the publisher stays quiet and we'd be testing
# min-load by accident — the exact failure mode this test exists to
# rule out.
_PREFIX_X_BODY = (
"Apricot bouquet cinnamon dewdrop elderflower fennel garlic "
"hibiscus indigo jasmine kumquat lavender mint nutmeg oregano "
"paprika quince rosemary saffron tarragon. "
)
PREFIX_X = (_PREFIX_X_BODY * 8).strip()
_PREFIX_Y_BODY = (
"Zephyr yellow xylophone wombat vortex umbrella thistle saffron "
"quartz peppermint orchid nightshade marigold lemongrass kale "
"juniper iris hyacinth gardenia foxglove. "
)
PREFIX_Y = (_PREFIX_Y_BODY * 8).strip()
_REQ_TOTAL_RE = re.compile(
r"^sgl_router_worker_requests_total\{([^}]*)\}\s+(\d+(?:\.\d+)?)\s*$"
)
_LABEL_RE = re.compile(r'(\w+)="([^"]*)"')
def _open_port() -> int:
with socket.socket() as sock:
sock.bind(("127.0.0.1", 0))
return sock.getsockname()[1]
@contextmanager
def _run(binary: Path, env: dict[str, str], log_path: Path):
with log_path.open("w") as log:
process = subprocess.Popen(
[str(binary)],
env={**os.environ, **env},
stdout=log,
stderr=subprocess.STDOUT,
start_new_session=True,
)
try:
yield process
finally:
process.terminate()
try:
process.wait(timeout=10)
except subprocess.TimeoutExpired:
process.kill()
process.wait(timeout=5)
def _wait_for_indexer(process: subprocess.Popen, port: int, log_path: Path) -> None:
deadline = time.time() + 10
while time.time() < deadline:
if process.poll() is not None:
raise RuntimeError(
f"KV Indexer exited during startup:\n{log_path.read_text()}"
)
try:
with socket.create_connection(("127.0.0.1", port), timeout=0.2):
return
except OSError:
time.sleep(0.1)
raise RuntimeError("timed out waiting for KV Indexer")
def _wait_for_bridge(process: subprocess.Popen, log_path: Path) -> None:
deadline = time.time() + 10
while time.time() < deadline:
output = log_path.read_text(errors="replace")
if "bridge session established" in output:
# ZMQ connect is asynchronous; let the subscription reach the PUB.
time.sleep(0.5)
return
if process.poll() is not None:
raise RuntimeError(f"KV Indexer Bridge exited during startup:\n{output}")
time.sleep(0.1)
raise RuntimeError(f"timed out waiting for KV Indexer Bridge:\n{output}")
def _dump_logs(logs: dict[str, Path]) -> None:
"""Print the tail of each Indexer/Bridge log so a routing failure is debuggable."""
for name, path in logs.items():
tail = path.read_text(errors="replace")[-4000:] if path.exists() else "<no log>"
print(f"\n----- {name} -----\n{tail}")
def _success_counts_by_worker(router_url: str) -> dict[str, int]:
"""Scrape ``/metrics`` and return ``{worker_url: success_count}``."""
r = httpx.get(f"{router_url}/metrics", timeout=5.0)
r.raise_for_status()
counts: dict[str, int] = {}
for line in r.text.splitlines():
m = _REQ_TOTAL_RE.match(line)
if not m:
continue
labels = dict(_LABEL_RE.findall(m.group(1)))
if labels.get("outcome") != "success":
continue
worker = labels.get("worker_url")
if not worker:
continue
try:
counts[worker] = counts.get(worker, 0) + int(float(m.group(2)))
except ValueError:
continue
return counts
def _send_chat(url: str, model_id: str, prompt: str) -> int:
"""POST one chat completion; return the HTTP status."""
r = httpx.post(
f"{url}/v1/chat/completions",
json={
"model": model_id,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": 4,
"stream": False,
},
timeout=60.0,
)
return r.status_code
def _direct_warm(worker_url: str, model_id: str, prefix: str) -> None:
"""Send one ``/v1/chat/completions`` request with ``prefix`` DIRECTLY to a worker.
The KV-event publisher emits ``BlockStored`` as the request's
prompt blocks commit to that worker's cache; routers subscribed to
the publisher receive the event and add ``(block_hash worker)``
entries to their ``HashTree``. The test then exercises those
entries by routing through the router.
Direct-warming (rather than going through a router) is the load-
bearing detail: routing through a router would itself choose which
worker to populate, so the two workers' HashTree state would no
longer be uniquely identifying.
Token alignment with the router the workers run with the model's
real chat template (no override), so the engine caches blocks keyed
on chat-templated tokens (role markers + content + generation prompt).
``cache_aware_zmq`` mirrors this: for a chat request on a model that
ships a chat template, it renders the same template and tokenizes the
result before hashing, so warm and route hash the same blocks.
"""
r = httpx.post(
f"{worker_url}/v1/chat/completions",
json={
"model": model_id,
"messages": [{"role": "user", "content": prefix}],
"max_tokens": 4,
"stream": False,
},
timeout=60.0,
)
assert r.status_code == 200, (
f"direct warm to {worker_url} failed: HTTP {r.status_code} {r.text!r}"
)
def _route_through(router_url: str, model_id: str, prompt: str) -> str:
"""Send one request through ``router_url``; return which worker handled it.
Computed by diffing the per-worker success-counter on ``/metrics``
around the call. Asserts exactly one worker absorbed the request
(no partial counts, no cancellation race).
"""
before = _success_counts_by_worker(router_url)
code = _send_chat(router_url, model_id, prompt)
assert code == 200, f"request to {router_url} failed: HTTP {code}"
after = _success_counts_by_worker(router_url)
deltas = {w: after.get(w, 0) - before.get(w, 0) for w in set(after) | set(before)}
winners = [w for w, d in deltas.items() if d > 0]
assert len(winners) == 1, (
f"expected exactly one worker delta on {router_url}, got {deltas}"
)
return winners[0]
@pytest.mark.real_gpu
@pytest.mark.slow
def test_routers_route_by_prefix_content(
router_binary,
gpu_allocator,
tmp_path,
):
"""Both the local ZMQ index and the external Indexer must route by content."""
spec = get_model_spec("qwen3-0.6b")
gpus = gpu_allocator.acquire(2)
indexer_port = _open_port()
indexer_endpoint = f"http://127.0.0.1:{indexer_port}"
indexer_binary = router_binary.parent / "kv-indexer-server"
bridge_binary = router_binary.parent / "kv-indexer-bridge"
logs = {
name: tmp_path / f"{name}.log" for name in ("indexer", "bridge-x", "bridge-y")
}
try:
with (
spawn_worker(
"qwen3-0.6b",
gpu_ids=[gpus[0]],
enable_kv_events=True,
) as worker_x,
spawn_worker(
"qwen3-0.6b",
gpu_ids=[gpus[1]],
enable_kv_events=True,
) as worker_y,
_run(
indexer_binary,
{"KV_INDEXER_LISTEN_ADDR": f"127.0.0.1:{indexer_port}"},
logs["indexer"],
) as indexer,
):
_wait_for_indexer(indexer, indexer_port, logs["indexer"])
worker_urls = [worker_x.url, worker_y.url]
def bridge_env(worker, worker_id: str) -> dict[str, str]:
assert worker.kv_events_endpoint is not None
return {
"KV_INDEXER_WORKER_ID": worker_id,
"KV_INDEXER_WORKER_ADDRESS": worker.url,
"KV_INDEXER_ENDPOINT": indexer_endpoint,
"SGLANG_KV_EVENT_ENDPOINT": worker.kv_events_endpoint.replace(
"*", "127.0.0.1"
),
"SGLANG_KV_EVENT_TOPIC": "kv",
}
with (
_run(
bridge_binary, bridge_env(worker_x, "worker-x"), logs["bridge-x"]
) as bridge_x,
_run(
bridge_binary, bridge_env(worker_y, "worker-y"), logs["bridge-y"]
) as bridge_y,
Gateway() as local,
Gateway() as external,
):
local.start_regular(
model_id=spec["model"],
tokenizer_path=spec["model"],
worker_urls=worker_urls,
policy="cache_aware_zmq",
timeout=120.0,
)
external.start_regular(
model_id=spec["model"],
tokenizer_path=spec["model"],
worker_urls=worker_urls,
policy="cache_aware_zmq",
kv_indexer_endpoint=indexer_endpoint,
timeout=120.0,
)
_wait_for_bridge(bridge_x, logs["bridge-x"])
_wait_for_bridge(bridge_y, logs["bridge-y"])
_direct_warm(worker_x.url, spec["model"], PREFIX_X)
_direct_warm(worker_y.url, spec["model"], PREFIX_Y)
time.sleep(2.0)
try:
for router, label in (
(local, "local-index"),
(external, "external-indexer"),
):
landed = _route_through(
router.base_url, spec["model"], PREFIX_X
)
assert landed == worker_x.url, (
f"router {label}: PREFIX_X must route to {worker_x.url}; landed on {landed}"
)
landed = _route_through(
router.base_url, spec["model"], PREFIX_Y
)
assert landed == worker_y.url, (
f"router {label}: PREFIX_Y must route to {worker_y.url}; landed on {landed}"
)
except Exception:
_dump_logs(logs)
raise
finally:
gpu_allocator.release(gpus)
@@ -186,7 +186,7 @@ class Gateway:
each worker's mode (plain) and any disaggregation
metadata are learned from ``/server_info``.
policy: Policy kind ``round_robin``, ``random``, ``power_of_two``,
or ``cache_aware_zmq``.
or ``cache_aware``.
kv_indexer_endpoint: Optional external KV Indexer gRPC endpoint.
timeout: How long to wait for ``/readyz`` before giving up.
"""
@@ -0,0 +1,35 @@
{
"ttft_slo_policy": "disabled",
"tps_slo_policy": "disabled",
"buckets": [
{
"id": "short",
"stage": "prefill",
"rank": 0,
"worker_ids": [
"http://127.0.0.1:31000",
"http://127.0.0.1:31001",
"http://127.0.0.1:31002",
"http://127.0.0.1:31003"
],
"min_extend_tokens": 0,
"max_extend_tokens": 2048,
"max_context_tokens": 32768,
"max_pending_prefill_tokens": 65536
},
{
"id": "long",
"stage": "prefill",
"rank": 1,
"worker_ids": [
"http://127.0.0.1:31004",
"http://127.0.0.1:31005",
"http://127.0.0.1:31006",
"http://127.0.0.1:31007"
],
"min_extend_tokens": 2049,
"max_context_tokens": 32768,
"max_pending_prefill_tokens": 65536
}
]
}
@@ -0,0 +1,794 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! HTTP contract for static P/D buckets.
//!
//! Buckets narrow the candidate domain before policy selection. Prefill SLO
//! profiles may override rank, while decode uses `input_tokens + max_tokens`.
use axum::body::Body;
use axum::http::{Request, StatusCode};
use sgl_kv_indexer::{PrefixIndex, PrefixIndexError, PrefixMatch, PrefixOutcome};
use sgl_router::config::{
ActiveLoadConfig, AffinityConfig, BucketConfig, BucketSpec, BucketStage, CacheAwareConfig,
CachePrefixProvider, Config, DiscoveryBackend, KvIndexerEndpointConfig, ModelConfig,
ObservabilityConfig, PolicyKind, ProxyConfig, ServerConfig, SessionAffinityMode,
SloBucketPolicy, StaticUrlsDiscoveryConfig,
};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::engine_load::{LoadStat, NativeCacheRankLoad};
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::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tower::ServiceExt;
fn bucket(id: &str, stage: BucketStage, rank: u32, worker_id: &str) -> BucketSpec {
BucketSpec {
id: id.into(),
stage,
rank,
worker_ids: vec![worker_id.into()],
min_extend_tokens: None,
max_extend_tokens: None,
min_sequence_tokens: None,
max_sequence_tokens: None,
max_context_tokens: Some(16_384),
ttft_p95_at_capacity_ms: None,
tps_p05_at_capacity: None,
max_pending_prefill_tokens: None,
}
}
fn build_app_context(
specs: Vec<WorkerSpec>,
bucket_config: BucketConfig,
policy: PolicyKind,
affinity: Option<AffinityConfig>,
) -> AppContext {
let config = Config {
server: ServerConfig {
host: "0".into(),
port: 0,
},
observability: ObservabilityConfig::default(),
model: ModelConfig {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy,
decode_policy: Default::default(),
bucket_config: Some(bucket_config),
circuit_breaker: None,
cache_aware: None,
sticky: None,
affinity,
fused: None,
eligibility: None,
},
discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()],
}),
proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(),
};
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&config).unwrap());
let registry = Arc::new(WorkerRegistry::default());
for spec in specs {
let _ = registry.add(spec);
}
let policies = Arc::new(build_registry_with_defaults(&config).unwrap());
let proxy = Arc::new(Proxy::new(Duration::from_secs(5)).unwrap());
AppContext::new(config, tokenizers, proxy, registry, policies)
}
fn build_ctx(
specs: Vec<WorkerSpec>,
bucket_config: BucketConfig,
policy: PolicyKind,
affinity: Option<AffinityConfig>,
) -> Arc<AppContext> {
Arc::new(build_app_context(specs, bucket_config, policy, affinity))
}
struct FakePrefixIndex {
address: Option<String>,
calls: AtomicUsize,
}
impl FakePrefixIndex {
fn matched(address: String) -> Arc<Self> {
Arc::new(Self {
address: Some(address),
calls: AtomicUsize::new(0),
})
}
fn no_signal() -> Arc<Self> {
Arc::new(Self {
address: None,
calls: AtomicUsize::new(0),
})
}
}
#[tonic::async_trait]
impl PrefixIndex for FakePrefixIndex {
async fn match_prefix(&self, hashes: Vec<i64>) -> Result<PrefixOutcome, PrefixIndexError> {
self.calls.fetch_add(1, Ordering::Relaxed);
let Some(address) = &self.address else {
return Ok(PrefixOutcome::Empty);
};
let matched_prefix_blocks =
u32::try_from(hashes.len().saturating_sub(1)).unwrap_or(u32::MAX);
Ok(PrefixOutcome::Matched {
matches: vec![PrefixMatch {
address: address.clone(),
matched_prefix_blocks,
worker_id: "fake-index-worker".into(),
}],
best_prefix_blocks: matched_prefix_blocks,
})
}
}
struct TwoPrefixIndex {
best_address: String,
lower_ranked_address: String,
}
impl TwoPrefixIndex {
fn new(best_address: String, lower_ranked_address: String) -> Arc<Self> {
Arc::new(Self {
best_address,
lower_ranked_address,
})
}
}
#[tonic::async_trait]
impl PrefixIndex for TwoPrefixIndex {
async fn match_prefix(&self, hashes: Vec<i64>) -> Result<PrefixOutcome, PrefixIndexError> {
let best_prefix_blocks = u32::try_from(hashes.len().saturating_sub(1)).unwrap_or(u32::MAX);
let lower_ranked_prefix_blocks = (best_prefix_blocks / 2).max(1);
Ok(PrefixOutcome::Matched {
matches: vec![
PrefixMatch {
address: self.best_address.clone(),
matched_prefix_blocks: best_prefix_blocks,
worker_id: "best-index-worker".into(),
},
PrefixMatch {
address: self.lower_ranked_address.clone(),
matched_prefix_blocks: lower_ranked_prefix_blocks,
worker_id: "lower-index-worker".into(),
},
],
best_prefix_blocks,
})
}
}
fn build_cache_ctx(
specs: Vec<WorkerSpec>,
bucket_config: BucketConfig,
prefix_index: Arc<dyn PrefixIndex>,
) -> Arc<AppContext> {
build_cache_ctx_with_affinity(
specs,
bucket_config,
prefix_index,
AffinityConfig::default(),
)
}
fn build_cache_ctx_with_affinity(
specs: Vec<WorkerSpec>,
bucket_config: BucketConfig,
prefix_index: Arc<dyn PrefixIndex>,
affinity: AffinityConfig,
) -> Arc<AppContext> {
let mut context =
build_app_context(specs, bucket_config, PolicyKind::CacheAware, Some(affinity));
context.config.model.cache_aware = Some(CacheAwareConfig {
prefix_provider: CachePrefixProvider::Indexer,
kv_indexer_endpoint: Some(KvIndexerEndpointConfig {
url: "http://fake-indexer".into(),
query_timeout_ms: 100,
query_max_inflight: 32,
}),
});
context.prefix_index = Some(prefix_index);
context.block_size_oracle.try_set(1).unwrap();
Arc::new(context)
}
fn worker_spec(id: &str, url: String, mode: WorkerMode) -> WorkerSpec {
WorkerSpec {
id: WorkerId(id.into()),
url,
mode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: (mode == WorkerMode::Prefill).then_some(8997),
}
}
fn set_native_load(
ctx: &AppContext,
worker_url: &str,
num_total_tokens: u64,
max_total_num_tokens: u64,
) {
ctx.engine_load.set(
worker_url,
0,
LoadStat {
num_running_reqs: 0,
num_waiting_reqs: 0,
num_tokens: num_total_tokens,
max_total_num_tokens,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 0,
num_total_tokens,
max_running_requests: 64,
total_prefill_uncached_tokens: 1,
total_prefill_busy_us: 1,
}),
},
Instant::now(),
);
}
fn chat_request(ttft_slo_ms: Option<u64>, max_tokens: Option<u64>) -> Request<Body> {
chat_request_with_content("bucket routing", ttft_slo_ms, max_tokens, None)
}
fn chat_request_with_content(
content: &str,
ttft_slo_ms: Option<u64>,
max_tokens: Option<u64>,
session_id: Option<&str>,
) -> Request<Body> {
let mut builder = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json");
if let Some(ttft_slo_ms) = ttft_slo_ms {
builder = builder.header("x-sgl-ttft-slo-ms", ttft_slo_ms.to_string());
}
if let Some(session_id) = session_id {
builder = builder.header("x-session-id", session_id);
}
builder
.body(Body::from(
serde_json::to_vec(&serde_json::json!({
"model": "tiny",
"messages": [{"role": "user", "content": content}],
"max_tokens": max_tokens,
}))
.unwrap(),
))
.unwrap()
}
async fn wait_for_prefill(mock: &crate::common::mock_worker::MockWorker) {
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if mock.captured.lock().unwrap().last_body.is_some() {
return;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.expect("selected prefill worker must receive the detached request");
}
async fn wait_for_prefill_body_containing(
mock: &crate::common::mock_worker::MockWorker,
expected: &str,
) -> Vec<u8> {
tokio::time::timeout(Duration::from_secs(2), async {
loop {
let captured = mock.captured.lock().unwrap().last_body.clone();
if let Some(body) = captured {
if String::from_utf8_lossy(&body).contains(expected) {
return body.to_vec();
}
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.expect("selected prefill worker must receive the expected request body")
}
#[tokio::test]
async fn prefill_slo_first_uses_eligible_ttft_bucket_before_lower_rank_bucket() {
let cheap = crate::common::mock_worker::MockWorker::start(vec![]).await;
let fast = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut cheap_bucket = bucket("p-cheap", BucketStage::Prefill, 10, "p-cheap");
cheap_bucket.ttft_p95_at_capacity_ms = Some(400);
let mut fast_bucket = bucket("p-fast", BucketStage::Prefill, 20, "p-fast");
fast_bucket.ttft_p95_at_capacity_ms = Some(100);
let bucket_config = BucketConfig {
buckets: vec![
cheap_bucket,
fast_bucket,
bucket("d-catch-all", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p-cheap", cheap.url.clone(), WorkerMode::Prefill),
worker_spec("p-fast", fast.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::PowerOfTwo,
None,
);
let response = build_router(ctx)
.oneshot(chat_request(Some(200), Some(16)))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
wait_for_prefill(&fast).await;
assert!(
cheap.captured.lock().unwrap().last_body.is_none(),
"lower-rank but TTFT-ineligible P Bucket must not be dispatched first"
);
}
#[tokio::test]
async fn prefill_tries_later_compatible_bucket_before_capacity_fallback() {
let full = crate::common::mock_worker::MockWorker::start(vec![]).await;
let available = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let bucket_config = BucketConfig {
buckets: vec![
bucket("p-full", BucketStage::Prefill, 10, "p-full"),
bucket("p-available", BucketStage::Prefill, 20, "p-available"),
bucket("d", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p-full", full.url.clone(), WorkerMode::Prefill),
worker_spec("p-available", available.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::PowerOfTwo,
None,
);
set_native_load(&ctx, &full.url, 100, 100);
set_native_load(&ctx, &available.url, 0, 10_000);
let response = build_router(ctx)
.oneshot(chat_request(None, Some(16)))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
wait_for_prefill(&available).await;
assert!(
full.captured.lock().unwrap().last_body.is_none(),
"capacity fallback must wait until all compatible prefill buckets are exhausted"
);
}
#[tokio::test]
async fn decode_bucket_uses_input_plus_requested_output_budget() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let short_decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let long_decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut short_bucket = bucket("d-short", BucketStage::Decode, 20, "d-short");
short_bucket.max_sequence_tokens = Some(1_024);
let mut long_bucket = bucket("d-long", BucketStage::Decode, 30, "d-long");
long_bucket.min_sequence_tokens = Some(1_025);
let bucket_config = BucketConfig {
buckets: vec![
bucket("p", BucketStage::Prefill, 10, "p"),
short_bucket,
long_bucket,
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p", prefill.url.clone(), WorkerMode::Prefill),
worker_spec("d-short", short_decode.url.clone(), WorkerMode::Decode),
worker_spec("d-long", long_decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::PowerOfTwo,
None,
);
let response = build_router(ctx)
.oneshot(chat_request(None, Some(2_000)))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("x-sgl-decode-url")
.and_then(|value| value.to_str().ok()),
Some(long_decode.url.as_str()),
"peak sequence length must exclude the short Decode Bucket"
);
assert!(
long_decode.captured.lock().unwrap().last_body.is_some(),
"the selected long Decode worker is awaited before the response"
);
assert!(
short_decode.captured.lock().unwrap().last_body.is_none(),
"the incompatible short Decode Bucket must not receive the request"
);
}
#[tokio::test]
async fn decode_tries_later_compatible_bucket_before_capacity_fallback() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let full = crate::common::mock_worker::MockWorker::start(vec![]).await;
let available = crate::common::mock_worker::MockWorker::start(vec![]).await;
let bucket_config = BucketConfig {
buckets: vec![
bucket("p", BucketStage::Prefill, 10, "p"),
bucket("d-full", BucketStage::Decode, 20, "d-full"),
bucket("d-available", BucketStage::Decode, 30, "d-available"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p", prefill.url.clone(), WorkerMode::Prefill),
worker_spec("d-full", full.url.clone(), WorkerMode::Decode),
worker_spec("d-available", available.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::PowerOfTwo,
None,
);
set_native_load(&ctx, &full.url, 100, 100);
set_native_load(&ctx, &available.url, 0, 10_000);
let response = build_router(ctx)
.oneshot(chat_request(None, Some(16)))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("x-sgl-decode-url")
.and_then(|value| value.to_str().ok()),
Some(available.url.as_str()),
"capacity fallback must wait until all compatible decode buckets are exhausted"
);
}
#[tokio::test]
async fn prefill_only_bucket_configuration_keeps_global_decode_routing() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let bucket_config = BucketConfig {
buckets: vec![bucket("p", BucketStage::Prefill, 10, "p")],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p", prefill.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::PowerOfTwo,
None,
);
let response = build_router(ctx)
.oneshot(chat_request(None, Some(16)))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("x-sgl-decode-url")
.and_then(|value| value.to_str().ok()),
Some(decode.url.as_str()),
"a Prefill-only Bucket rollout must retain the Step 1 global Decode domain"
);
}
#[tokio::test]
async fn global_rebind_session_affinity_can_keep_a_cross_length_bucket_primary() {
let short = crate::common::mock_worker::MockWorker::start(vec![]).await;
let long = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut short_bucket = bucket("p-short", BucketStage::Prefill, 10, "p-short");
short_bucket.max_extend_tokens = Some(256);
short_bucket.max_context_tokens = Some(16_384);
short_bucket.ttft_p95_at_capacity_ms = Some(80);
let mut long_bucket = bucket("p-long", BucketStage::Prefill, 20, "p-long");
long_bucket.min_extend_tokens = Some(257);
long_bucket.max_context_tokens = Some(16_384);
long_bucket.ttft_p95_at_capacity_ms = Some(300);
let bucket_config = BucketConfig {
buckets: vec![
short_bucket,
long_bucket,
bucket("d-catch-all", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::SloFirst,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p-short", short.url.clone(), WorkerMode::Prefill),
worker_spec("p-long", long.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::SessionAware,
Some(AffinityConfig {
session_affinity_mode: SessionAffinityMode::GlobalRebind,
..Default::default()
}),
);
let app = build_router(ctx);
let first = app
.clone()
.oneshot(chat_request_with_content(
"short",
Some(120),
Some(8),
Some("s-1"),
))
.await
.unwrap();
assert_eq!(first.status(), StatusCode::OK);
wait_for_prefill(&short).await;
let long_content = "length ".repeat(128);
let second = app
.oneshot(chat_request_with_content(
&long_content,
Some(120),
Some(8),
Some("s-1"),
))
.await
.unwrap();
assert_eq!(second.status(), StatusCode::OK);
let short_body = wait_for_prefill_body_containing(&short, &long_content).await;
assert!(
String::from_utf8_lossy(&short_body).contains(&long_content),
"the second, long request must retain the existing cross-Bucket session primary"
);
assert!(
long.captured.lock().unwrap().last_body.is_none(),
"target length Bucket is skipped only because the primary's own Hard TTFT profile is eligible"
);
}
#[tokio::test]
async fn global_preserve_establishes_then_reuses_a_new_assignment() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let bucket_config = BucketConfig {
buckets: vec![
bucket("p", BucketStage::Prefill, 10, "p"),
bucket("d", BucketStage::Decode, 20, "d"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let ctx = build_ctx(
vec![
worker_spec("p", prefill.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
PolicyKind::SessionAware,
Some(AffinityConfig {
session_affinity_mode: SessionAffinityMode::GlobalPreserve,
..Default::default()
}),
);
let app = build_router(Arc::clone(&ctx));
for content in ["first global request", "second global request"] {
let response = app
.clone()
.oneshot(chat_request_with_content(
content,
None,
Some(8),
Some("global-session"),
))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
let metrics = ctx.metrics.render();
assert!(
metrics.contains(
r#"sgl_router_policy_decisions_total{policy="session_aware",reason="assigned"} 1"#
),
"the first global-preserve request must establish an assignment: {metrics}"
);
assert!(
metrics.contains(
r#"sgl_router_policy_decisions_total{policy="session_aware",reason="session_primary"} 1"#
),
"the second global-preserve request must reuse the assignment: {metrics}"
);
}
#[tokio::test]
async fn cache_winner_uses_target_uncached_work_before_prompt_length_bucket() {
let short = crate::common::mock_worker::MockWorker::start(vec![]).await;
let long = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut short_bucket = bucket("p-short", BucketStage::Prefill, 10, "p-short");
short_bucket.max_extend_tokens = Some(8);
let mut long_bucket = bucket("p-long", BucketStage::Prefill, 20, "p-long");
long_bucket.min_extend_tokens = Some(9);
let bucket_config = BucketConfig {
buckets: vec![
short_bucket,
long_bucket,
bucket("d-catch-all", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let index = FakePrefixIndex::matched(short.url.clone());
let prefix_index: Arc<dyn PrefixIndex> = index.clone();
let ctx = build_cache_ctx(
vec![
worker_spec("p-short", short.url.clone(), WorkerMode::Prefill),
worker_spec("p-long", long.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
prefix_index,
);
let content = "cached-prefix ".repeat(128);
let response = build_router(ctx)
.oneshot(chat_request_with_content(&content, None, Some(8), None))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
wait_for_prefill(&short).await;
assert!(
long.captured.lock().unwrap().last_body.is_none(),
"a cache winner with small target-specific uncached work must not be replaced by the full-length Bucket"
);
assert_eq!(
index.calls.load(Ordering::Relaxed),
1,
"the async Indexer query must run once at ingress, not once per Bucket"
);
}
#[tokio::test]
async fn cache_candidate_bucket_binding_happens_before_candidate_limit() {
let best = crate::common::mock_worker::MockWorker::start(vec![]).await;
let lower_ranked = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut best_bucket = bucket("p-best", BucketStage::Prefill, 10, "p-best");
best_bucket.min_extend_tokens = Some(32);
let mut lower_ranked_bucket = bucket("p-lower", BucketStage::Prefill, 20, "p-lower");
lower_ranked_bucket.min_extend_tokens = Some(32);
let bucket_config = BucketConfig {
buckets: vec![
best_bucket,
lower_ranked_bucket,
bucket("d-catch-all", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let index: Arc<dyn PrefixIndex> =
TwoPrefixIndex::new(best.url.clone(), lower_ranked.url.clone());
let ctx = build_cache_ctx_with_affinity(
vec![
worker_spec("p-best", best.url.clone(), WorkerMode::Prefill),
worker_spec("p-lower", lower_ranked.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
index,
AffinityConfig {
cache_candidate_min_workers: 1,
cache_candidate_ratio: 0.0,
cache_candidate_max_workers: 1,
..AffinityConfig::default()
},
);
let content = "cached bucket candidate ".repeat(256);
let response = build_router(Arc::clone(&ctx))
.oneshot(chat_request_with_content(&content, None, Some(8), None))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
wait_for_prefill(&lower_ranked).await;
assert!(
best.captured.lock().unwrap().last_body.is_none(),
"the top Indexer hit is Bucket-incompatible and must not consume K=1"
);
assert!(
ctx.metrics.render().contains(
r#"sgl_router_policy_decisions_total{policy="cache_aware",reason="cache_candidate"} 1"#
),
"the compatible lower-ranked cache holder must remain a cache candidate"
);
}
#[tokio::test]
async fn cache_no_signal_restarts_normal_prompt_length_bucket_fallback() {
let short = crate::common::mock_worker::MockWorker::start(vec![]).await;
let long = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start(vec![]).await;
let mut short_bucket = bucket("p-short", BucketStage::Prefill, 10, "p-short");
short_bucket.max_extend_tokens = Some(8);
let mut long_bucket = bucket("p-long", BucketStage::Prefill, 20, "p-long");
long_bucket.min_extend_tokens = Some(9);
let bucket_config = BucketConfig {
buckets: vec![
short_bucket,
long_bucket,
bucket("d-catch-all", BucketStage::Decode, 30, "d"),
],
ttft_slo_policy: SloBucketPolicy::Disabled,
tps_slo_policy: SloBucketPolicy::Disabled,
};
let index = FakePrefixIndex::no_signal();
let prefix_index: Arc<dyn PrefixIndex> = index.clone();
let ctx = build_cache_ctx(
vec![
worker_spec("p-short", short.url.clone(), WorkerMode::Prefill),
worker_spec("p-long", long.url.clone(), WorkerMode::Prefill),
worker_spec("d", decode.url.clone(), WorkerMode::Decode),
],
bucket_config,
prefix_index,
);
let content = "uncached-prompt ".repeat(128);
let response = build_router(ctx)
.oneshot(chat_request_with_content(&content, None, Some(8), None))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
wait_for_prefill(&long).await;
assert!(
short.captured.lock().unwrap().last_body.is_none(),
"without a cache winner the request must restart the normal full-input Bucket path"
);
assert_eq!(index.calls.load(Ordering::Relaxed), 1);
}
@@ -12,20 +12,11 @@
//! doesn't render tool schemas, so its ids would diverge from the engine).
//! * A request with multimodal (array) content → `input_ids` omitted (a text
//! tokenizer can't represent image content).
//!
//! The model id contains `deepseek-v4` so the tokenizer registry auto-attaches
//! the built-in V4 chat encoder — the engine-equivalent path — without a
//! template fixture.
use axum::body::Body;
use axum::http::{Request, StatusCode};
use serde_json::{json, Value};
use sgl_router::config::{
ActiveLoadConfig, CacheAwareConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::engine_load::EngineLoadTable;
use sgl_router::policies::factory::build_registry;
use sgl_router::policies::kv_events::{BlockSizeOracle, HashTree};
use sgl_router::proxy::Proxy;
@@ -37,36 +28,9 @@ use std::sync::Arc;
use std::time::Duration;
use tower::ServiceExt;
use crate::common::cache_aware_fixture::{config, MODEL};
use crate::common::mock_worker::MockWorker;
const MODEL: &str = "deepseek-v4-tiny";
fn config() -> Config {
Config {
server: ServerConfig {
host: "0".into(),
port: 0,
},
observability: ObservabilityConfig::default(),
model: ModelConfig {
id: MODEL.into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::CacheAwareZmq,
circuit_breaker: None,
cache_aware: Some(CacheAwareConfig::default()),
sticky: None,
affinity: None,
fused: None,
eligibility: None,
},
discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()],
}),
proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(),
}
}
fn build_ctx(url: String) -> Arc<AppContext> {
let cfg = config();
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
@@ -82,18 +46,9 @@ fn build_ctx(url: String) -> Arc<AppContext> {
model_ids: vec![ModelId(MODEL.into())],
bootstrap_port: None,
});
// Use the real loaded tokenizers (not the empty-registry test default) so
// the cache-aware policy can tokenize at ingress.
let policies = Arc::new(
build_registry(
&cfg,
Arc::new(HashTree::new()),
Arc::clone(&tokenizers),
BlockSizeOracle::new(),
EngineLoadTable::new(),
)
.unwrap(),
);
// Use the configured tokenizer so the chat path can emit input_ids.
let policies =
Arc::new(build_registry(&cfg, Arc::new(HashTree::new()), BlockSizeOracle::new()).unwrap());
let proxy = Arc::new(Proxy::new(Duration::from_secs(5)).unwrap());
Arc::new(AppContext::new(cfg, tokenizers, proxy, registry, policies))
}
@@ -34,6 +34,8 @@ fn config_for(_worker_url: &str) -> Config {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -13,7 +13,7 @@ use sgl_router::config::{
pub const MODEL: &str = "deepseek-v4-tiny";
/// A single-model `cache_aware_zmq` router. Discovery is a placeholder because
/// A single-model native `cache_aware` router. Discovery is a placeholder because
/// every caller installs its own `WorkerRegistry`.
pub fn config() -> Config {
Config {
@@ -25,11 +25,13 @@ pub fn config() -> Config {
model: ModelConfig {
id: MODEL.into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::CacheAwareZmq,
policy: PolicyKind::CacheAware,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: Some(CacheAwareConfig::default()),
sticky: None,
affinity: None,
sticky: None,
fused: None,
eligibility: None,
},
@@ -16,8 +16,8 @@ use sgl_kv_indexer::pb::{
use sgl_kv_indexer::{
server_builder, GrpcPrefixIndex, InMemoryKvIndexerBackend, KvIndexerService, PrefixIndexConfig,
};
use sgl_router::config::{AffinityConfig, CachePrefixProvider, PolicyKind};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::engine_load::EngineLoadTable;
use sgl_router::policies::factory::build_registry;
use sgl_router::policies::kv_events::{compute_block_hashes, BlockSizeOracle, HashTree};
use sgl_router::policies::request_tokens_for;
@@ -36,7 +36,20 @@ use crate::common::mock_worker::MockWorker;
async fn external_indexer_routes_to_the_cached_worker() {
let cached = MockWorker::start(vec![]).await;
let uncached = MockWorker::start(vec![]).await;
let cfg = config();
let mut cfg = config();
cfg.model.policy = PolicyKind::CacheAware;
cfg.model
.cache_aware
.as_mut()
.expect("fixture includes cache-aware configuration")
.prefix_provider = CachePrefixProvider::Indexer;
cfg.model.affinity = Some(AffinityConfig {
cache_affinity_min_matched_tokens: Some(0),
cache_candidate_min_workers: 1,
cache_candidate_ratio: 1.0,
cache_candidate_max_workers: 1,
..Default::default()
});
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let body = json!({
"model": MODEL,
@@ -68,6 +81,7 @@ async fn external_indexer_routes_to_the_cached_worker() {
hashes: hashes.clone(),
component_masks: Vec::new(),
block_sizes: Vec::new(),
parent_block_hash: None,
}],
worker_address: cached.url.clone(),
cache_spec: None,
@@ -89,16 +103,8 @@ async fn external_indexer_routes_to_the_cached_worker() {
}
let oracle = BlockSizeOracle::new();
oracle.try_set(1).unwrap();
let policies = Arc::new(
build_registry(
&cfg,
Arc::new(HashTree::new()),
Arc::clone(&tokenizers),
Arc::clone(&oracle),
EngineLoadTable::new(),
)
.unwrap(),
);
let policies =
Arc::new(build_registry(&cfg, Arc::new(HashTree::new()), Arc::clone(&oracle)).unwrap());
let mut ctx = AppContext::new(
cfg,
tokenizers,
@@ -35,6 +35,8 @@ async fn failover_when_one_worker_dies() {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: Some(CircuitBreakerConfig {
threshold: std::num::NonZeroU32::new(1).unwrap(), // open after first failure
cool_down_secs: 30,
@@ -14,7 +14,7 @@
//! guards, or in the SSE pump's `tx.send().await` race — all of which
//! would be silently skipped by a synthetic-handler test.
use bytes::Bytes;
use futures::future::join_all;
use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
@@ -44,6 +44,8 @@ fn build_ctx_with_worker(worker_url: &str) -> Arc<AppContext> {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -130,12 +132,11 @@ async fn shutdown_drains_100_inflight_streaming_chat_completions() {
}))
.unwrap();
let mut handles = Vec::with_capacity(N);
for i in 0..N {
let responses = join_all((0..N).map(|i| {
let c = client.clone();
let u = url.clone();
let b = body.clone();
handles.push(tokio::spawn(async move {
async move {
let resp = c
.post(&u)
.header("content-type", "application/json")
@@ -146,33 +147,34 @@ async fn shutdown_drains_100_inflight_streaming_chat_completions() {
if !resp.status().is_success() {
return Err(format!("client {i} non-2xx: {}", resp.status()));
}
let bytes: Bytes = resp
.bytes()
.await
.map_err(|e| format!("client {i} body: {e}"))?;
Ok::<Bytes, String>(bytes)
}));
}
Ok::<_, String>((i, resp))
}
}))
.await;
let responses: Vec<_> = responses
.into_iter()
.collect::<Result<_, _>>()
.expect("every client received response headers before shutdown");
// 4. Let every request grab a connection and start receiving data.
// 100 ms is past the first chunk delay (60 ms) for every stream
// but well before the last chunk fires.
tokio::time::sleep(Duration::from_millis(100)).await;
// 5. Trigger shutdown. axum stops accepting new connections but
// MUST drain the 100 already-attached streams.
// 4. Each response header confirms that its request is in flight. Trigger
// shutdown only after the full cohort connects, then verify that Axum
// drains all 100 existing streams.
let started = Instant::now();
shutdown_tx.send(()).unwrap();
// 6. Every in-flight request must complete with a `[DONE]` terminator
// 5. Every in-flight request must complete with a `[DONE]` terminator
// — proving the stream was NOT truncated by shutdown.
let mut bytes_total: usize = 0;
let mut done_count: usize = 0;
for h in handles {
let result = h
for result in join_all(responses.into_iter().map(|(i, response)| async move {
response
.bytes()
.await
.expect("client task panicked")
.expect("client completed");
.map_err(|e| format!("client {i} body: {e}"))
}))
.await
{
let result = result.expect("client body completed");
bytes_total += result.len();
let body_str = String::from_utf8_lossy(&result);
if body_str.contains("data: [DONE]") {
@@ -31,6 +31,8 @@ async fn forwards_whitelisted_headers_strips_others() {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -10,6 +10,7 @@
mod common;
mod bucket_routing;
mod cache_aware_input_ids;
mod chat_routing;
mod external_indexer_routing;
@@ -18,6 +19,7 @@ mod graceful_shutdown;
mod header_forwarding;
mod pd_bootstrap_injection;
mod pd_pool_isolation;
mod radix_tree_routing;
mod roundrobin_input_ids;
mod shared_prefill_admission;
mod sticky_input_ids;
@@ -32,7 +32,7 @@ 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 std::time::{Duration, Instant};
use tower::ServiceExt;
fn config() -> Config {
@@ -46,6 +46,8 @@ fn config() -> Config {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -98,7 +100,7 @@ async fn await_captured_body(
timeout: Duration,
label: &str,
) -> Bytes {
let start = std::time::Instant::now();
let start = Instant::now();
loop {
// Release the `std::sync::Mutex` guard before the sleep.await
// (clippy: await_holding_lock).
@@ -189,6 +191,43 @@ async fn pd_mode_chat_injects_bootstrap_fields_into_both_bodies() {
assert_eq!(bootstrap_port(&dj), Some(8997));
}
#[tokio::test]
async fn round_robin_pd_prefill_does_not_track_dispatch_timestamps() {
let prefill =
crate::common::mock_worker::MockWorker::start_hanging(Duration::from_millis(200)).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 prefill_worker = ctx
.registry
.workers_for(&ModelId("tiny".into()))
.into_iter()
.find(|worker| worker.id.0 == "p1")
.expect("prefill worker is registered");
let cutoff = Instant::now() - Duration::from_secs(1);
let request = tokio::spawn(build_router(Arc::clone(&ctx)).oneshot(chat_request()));
await_captured_body(&prefill, Duration::from_secs(2), "prefill").await;
assert_eq!(prefill_worker.active_load(), 1);
assert_eq!(prefill_worker.slots_acquired_since(cutoff), 0);
assert_eq!(request.await.unwrap().unwrap().status(), StatusCode::OK);
}
/// Plain-mode (non-PD) requests do NOT carry any `bootstrap_*` field.
/// The injection step is gated on `worker.mode() == Prefill`; plain
/// workers serve the chat route directly without disagg bootstrapping.
@@ -45,6 +45,8 @@ fn config() -> Config {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -205,31 +207,19 @@ async fn pd_mode_chat_dispatch_fans_to_both_prefill_and_decode() {
assert!(!prefill_body.is_empty());
}
/// Task C: PD-mode chat request carries an `x-sgl-decode-url` header
/// pointing at the host-affinity decode peer. With two prefill workers
/// on different hosts and a decode worker on each, the affinity helper
/// MUST pick the decode peer co-located with the chosen prefill.
///
/// Round-robin will select prefill workers deterministically (alphabetic
/// dashmap order is not guaranteed; the test fires several requests so
/// at least one lands on each prefill, and asserts the per-host pairing
/// holds across all of them).
/// 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_decode_affinity_header() {
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 always bind to `127.0.0.1`, so every worker
// shares the same host string and the affinity helper's
// same-host branch is moot here — the helper still returns a
// decode peer via the load-tiebreak fallback. The unit tests in
// `policies::registry::tests::decoder_picks_same_host_when_available`
// carry the real burden of pinning the host-affinity rules; this
// integration test only asserts the wiring is in place (the
// `x-sgl-decode-url` header IS set on PD requests, and the
// value is one of the registered decode worker URLs).
// 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()),
@@ -268,9 +258,8 @@ async fn pd_mode_chat_dispatch_sets_decode_affinity_header() {
assert_eq!(res.status(), StatusCode::OK);
}
// Every request that hit a prefill mock MUST carry the decode-hint
// header. The header value MUST be one of the two registered
// decode worker URLs.
// 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();
@@ -346,9 +335,9 @@ async fn pd_mode_prefill_only_returns_no_decode_workers_available() {
}
/// PD-mode chat response carries `x-sgl-decode-url` so external tests
/// can observe decode affinity end-to-end (without sniffing the proxy
/// 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_decode_affinity_header`.
/// 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;
@@ -0,0 +1,94 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use std::sync::Arc;
use std::time::Duration;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use serde_json::json;
use sgl_router::config::{AffinityConfig, CachePrefixProvider, PolicyKind};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry;
use sgl_router::policies::kv_events::{
compute_block_hashes, BlockSizeOracle, HashTree, KvWorkerId,
};
use sgl_router::policies::prefix_provider::RadixTreePrefixProvider;
use sgl_router::policies::request_tokens_for;
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 tower::ServiceExt;
use crate::common::cache_aware_fixture::{config, MODEL};
use crate::common::mock_worker::MockWorker;
#[tokio::test]
async fn radix_tree_routes_cache_aware_request_to_cached_worker() {
let cached = MockWorker::start(vec![]).await;
let uncached = MockWorker::start(vec![]).await;
let mut cfg = config();
cfg.model.policy = PolicyKind::CacheAware;
cfg.model.cache_aware.as_mut().unwrap().prefix_provider = CachePrefixProvider::RadixTree;
cfg.model.affinity = Some(AffinityConfig {
cache_affinity_min_matched_tokens: Some(0),
cache_candidate_min_workers: 1,
cache_candidate_ratio: 1.0,
cache_candidate_max_workers: 1,
..Default::default()
});
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let body = json!({
"model": MODEL,
"messages": [{"role": "user", "content": "local radix cache hit"}],
});
let tokens = request_tokens_for(&tokenizers, &ModelId(MODEL.into()), &body)
.expect("test prompt tokenizes");
let hashes = compute_block_hashes(&tokens.ids, 1);
assert!(!hashes.is_empty());
let tree = Arc::new(HashTree::new());
tree.insert(&KvWorkerId::new(cached.url.clone(), 0), None, &hashes);
let registry = Arc::new(WorkerRegistry::default());
for url in [&cached.url, &uncached.url] {
registry
.add(WorkerSpec {
id: WorkerId(url.clone()),
url: url.clone(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId(MODEL.into())],
bootstrap_port: None,
})
.unwrap();
}
let oracle = BlockSizeOracle::new();
oracle.try_set(1).unwrap();
let policies = Arc::new(build_registry(&cfg, Arc::clone(&tree), Arc::clone(&oracle)).unwrap());
let mut ctx = AppContext::new(
cfg,
tokenizers,
Arc::new(Proxy::new(Duration::from_secs(5)).unwrap()),
registry,
policies,
);
ctx.radix_tree_prefix_provider = Some(RadixTreePrefixProvider::new(tree, Arc::clone(&oracle)));
ctx.block_size_oracle = oracle;
let response = build_router(Arc::new(ctx))
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(cached.captured.lock().unwrap().last_body.is_some());
assert!(uncached.captured.lock().unwrap().last_body.is_none());
}
@@ -42,6 +42,8 @@ fn config() -> Config {
id: MODEL.into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -11,7 +11,7 @@ use sgl_router::config::{
ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::engine_load::LoadStat;
use sgl_router::policies::engine_load::{LoadStat, NativeCacheRankLoad};
use sgl_router::policies::{
CacheCandidate, CacheCandidateProposal, Policy, PolicyRegistry, PrefillProposal, ProposalKind,
SelectionContext, SelectionProposal,
@@ -128,6 +128,7 @@ impl Policy for CacheCandidatesPolicy {
max_pending_prefill_tokens: None,
}],
cache_switch_margin_tokens: 0,
..Default::default()
}))
}
@@ -147,6 +148,8 @@ fn config(policy: PolicyKind) -> Config {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,
@@ -286,17 +289,30 @@ async fn chat_commits_the_admitted_prefill_backup() {
})
})
.await;
let now = Instant::now();
let native_load = |total_prefill_uncached_tokens, total_prefill_busy_us| LoadStat {
num_running_reqs: 1,
num_waiting_reqs: 0,
num_tokens: 100,
max_total_num_tokens: 100,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 0,
num_total_tokens: 100,
max_running_requests: 16,
total_prefill_uncached_tokens,
total_prefill_busy_us,
}),
};
fixture.ctx.engine_load.set(
&fixture.workers[0].url,
0,
LoadStat {
num_running_reqs: 1,
num_waiting_reqs: 0,
num_tokens: 100,
max_total_num_tokens: 100,
},
Instant::now(),
native_load(1, 1),
now - Duration::from_secs(1),
);
fixture
.ctx
.engine_load
.set(&fixture.workers[0].url, 0, native_load(2, 2), now);
assert_eq!(send_chat(&fixture.ctx).await, StatusCode::OK);
assert!(fixture.backends[0]
@@ -339,6 +355,13 @@ async fn capacity_exhaustion_does_not_return_503() {
num_waiting_reqs: 0,
num_tokens: 100,
max_total_num_tokens: 100,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 0,
num_total_tokens: 100,
max_running_requests: 16,
total_prefill_uncached_tokens: 1,
total_prefill_busy_us: 1,
}),
},
Instant::now(),
);
@@ -415,6 +438,13 @@ async fn chat_records_cache_candidates_exhausted() {
num_waiting_reqs: 0,
num_tokens: 100,
max_total_num_tokens: 100,
native_cache: Some(NativeCacheRankLoad {
num_waiting_uncached_tokens: 0,
num_total_tokens: 100,
max_running_requests: 16,
total_prefill_uncached_tokens: 1,
total_prefill_busy_us: 1,
}),
},
Instant::now(),
);
@@ -54,6 +54,8 @@ fn config() -> Config {
id: MODEL.into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::Sticky,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
// Push eviction far out so the background sweeper never fires
@@ -42,6 +42,8 @@ fn build_sticky_ctx(header_name: &str, worker_urls: &[String]) -> Arc<AppContext
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::Sticky,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: Some(StickyConfig {
@@ -38,6 +38,8 @@ fn config(_worker_url: &str) -> Config {
id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin,
decode_policy: Default::default(),
bucket_config: None,
circuit_breaker: None,
cache_aware: None,
sticky: None,