[router] Configure experimental sgl-router via CLI flags instead of a config file (#27073)

Signed-off-by: Kangyan Zhou <zky314343421@gmail.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-06-05 10:02:10 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 631db6c757
commit bcf89928b4
43 changed files with 2355 additions and 1237 deletions
+5 -6
View File
@@ -37,10 +37,11 @@ reqwest = { version = "0.12", features = ["stream", "json", "rustls-tls"], defau
# Serialization # Serialization
serde = { version = "1", features = ["derive"] } serde = { version = "1", features = ["derive"] }
serde_json = { version = "1", features = ["preserve_order"] } serde_json = { version = "1", features = ["preserve_order"] }
# `humantime-serde` lets `WorkerConfig.request_timeout` accept human-readable
# durations like `"60s"` / `"500ms"` / `"2m"` in YAML / TOML, rather than # Tokenizer auto-download from HuggingFace when --tokenizer-path is omitted.
# forcing operators to write raw milliseconds. # Sync (`ureq`) API only — it runs once at startup; `ureq` is on rustls, so
humantime-serde = "1" # this pulls no openssl/native-tls (matching reqwest's rustls-tls above).
hf-hub = { version = "0.4", default-features = false, features = ["ureq"] }
# Utilities # Utilities
anyhow = "1" anyhow = "1"
@@ -53,8 +54,6 @@ bytes = "1"
rand = "0.8" rand = "0.8"
tokio-stream = "0.1" tokio-stream = "0.1"
dashmap = "6" dashmap = "6"
serde_yaml = "0.9"
toml = "0.8"
kube = { version = "0.96", features = ["runtime", "derive"] } kube = { version = "0.96", features = ["runtime", "derive"] }
k8s-openapi = { version = "0.23", features = ["v1_31"] } k8s-openapi = { version = "0.23", features = ["v1_31"] }
tokio-util = "0.7" tokio-util = "0.7"
+40 -5
View File
@@ -2,11 +2,11 @@
Slim, KV-aware, OpenAI-compatible router for SGLang workers. Slim, KV-aware, OpenAI-compatible router for SGLang workers.
**Status:** functional single-worker HTTP proxy. Exposes `/v1/tokenize`, Serves a single model and routes across its workers. Exposes
`/v1/detokenize`, `/v1/models`, `/v1/chat/completions` (buffered and SSE), `/v1/tokenize`, `/v1/detokenize`, `/v1/models`, `/v1/chat/completions`
plus `/healthz` / `/readyz`. Forwards to one configured worker via reqwest; (buffered and SSE), plus `/healthz` / `/readyz` and `/metrics`. Worker
parity-tested against `transformers.AutoTokenizer`. Multi-worker routing, pools come from either a static URL list or Kubernetes EndpointSlice
service discovery, and observability still pending. discovery.
## Building ## Building
@@ -15,6 +15,41 @@ cd experimental/sgl-router
cargo build --release cargo build --release
``` ```
## Running
The router is configured entirely through CLI flags (run
`sgl-router --help` for the full list). It serves exactly one model, so
`--model-id` is required, along with exactly one discovery backend.
`--tokenizer-path` is optional: give it a local `tokenizer.json` path or a
HuggingFace repo id, and when omitted the router downloads the tokenizer
for `--model-id` from HuggingFace (honoring `HF_TOKEN` / `HF_HOME`).
Static worker list:
```bash
sgl-router \
--host 0.0.0.0 --port 30000 \
--model-id qwen3 \
--tokenizer-path /models/qwen3/tokenizer.json \
--worker-urls http://10.0.0.1:30000 http://10.0.0.2:30000
```
Kubernetes EndpointSlice discovery:
```bash
sgl-router \
--host 0.0.0.0 --port 30000 \
--model-id qwen3 \
--tokenizer-path /models/qwen3/tokenizer.json \
--service-discovery \
--service-discovery-namespace prod \
--selector app=engines-qwen3
```
Omit `--service-discovery-namespace` to watch all namespaces (requires
cluster-wide RBAC). For prefill/decode disaggregation, replace `--selector`
with `--prefill-selector` and `--decode-selector`.
## License ## License
Apache-2.0. Apache-2.0.
+711
View File
@@ -0,0 +1,711 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! Command-line interface. The router is configured entirely through
//! flags — there is no config file. [`Cli::into_config`] resolves the
//! flags into a validated [`Config`].
use anyhow::{anyhow, Result};
use clap::Parser;
use std::num::NonZeroU32;
use crate::config::{
default_cb_cool_down, default_proxy_request_timeout_secs, default_stale_request_timeout_secs,
resolve_mode, ActiveLoadConfig, CacheAwareConfig, CircuitBreakerConfig, Config,
DiscoveryBackend, K8sDiscoveryConfig, LogFormat, ModelConfig, ObservabilityConfig, PolicyKind,
ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
};
/// `sgl-router` — slim KV-aware OpenAI-compatible router for SGLang workers.
///
/// Discovery is mutually exclusive: pass `--worker-urls` for a static
/// worker list, or `--service-discovery` for Kubernetes EndpointSlice
/// discovery — exactly one is required.
#[derive(Parser, Debug)]
#[command(
name = "sgl-router",
version,
about = "Slim KV-aware OpenAI-compatible router for SGLang workers"
)]
pub struct Cli {
// ---- server ----
/// Address to bind the HTTP server to.
#[arg(long, default_value = "127.0.0.1")]
pub host: String,
/// Port to bind the HTTP server to.
#[arg(long, default_value_t = 30000)]
pub port: u16,
// ---- model (exactly one) ----
/// Model id this router serves (the OpenAI `model` field).
#[arg(long)]
pub model_id: String,
/// Tokenizer source: a local `tokenizer.json` path, or a HuggingFace
/// repo id to download from. When omitted, falls back to `--model-id`
/// as the repo id (download honors `HF_TOKEN` / `HF_HOME`).
#[arg(long)]
pub tokenizer_path: Option<String>,
/// Routing policy.
#[arg(long, value_enum, default_value = "round_robin")]
pub policy: PolicyKind,
// ---- circuit breaker (opt-in via --cb-threshold) ----
/// Consecutive upstream failures before the circuit breaker opens.
/// Setting this enables the circuit breaker; `0` is rejected.
#[arg(long)]
pub cb_threshold: Option<NonZeroU32>,
/// Circuit-breaker cool-down in seconds. Only meaningful with
/// `--cb-threshold`; defaults to 30 when the breaker is enabled.
#[arg(long)]
pub cb_cool_down_secs: Option<u64>,
// ---- cache-aware-zmq tuning (only used by that policy) ----
/// 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>,
// ---- discovery: static ----
/// Static worker URLs (space-separated or repeated). Mutually
/// exclusive with `--service-discovery`.
#[arg(long, num_args = 1..)]
pub worker_urls: Vec<String>,
// ---- discovery: kubernetes ----
/// Enable Kubernetes EndpointSlice discovery.
#[arg(long)]
pub service_discovery: bool,
/// Namespace to watch. Unset/empty watches all namespaces (requires
/// cluster-wide RBAC).
#[arg(long)]
pub service_discovery_namespace: Option<String>,
/// Plain-mode label selector terms, e.g. `app=engines-qwen3`
/// (space-separated or repeated `key=value`, AND-joined). Mutually
/// exclusive with the prefill/decode selectors.
#[arg(long, num_args = 1..)]
pub selector: Vec<String>,
/// PD-mode prefill label selector terms. Requires `--decode-selector`.
#[arg(long, num_args = 1..)]
pub prefill_selector: Vec<String>,
/// PD-mode decode label selector terms. Requires `--prefill-selector`.
#[arg(long, num_args = 1..)]
pub decode_selector: Vec<String>,
// ---- proxy / active-load ----
/// Per-request upstream timeout in seconds.
#[arg(long, default_value_t = default_proxy_request_timeout_secs())]
pub request_timeout_secs: u64,
/// Max lifetime of an in-flight request entry before the janitor
/// reaps it (returns 504 `stale_request_expired`).
#[arg(long, default_value_t = default_stale_request_timeout_secs())]
pub stale_request_timeout_secs: u64,
// ---- observability ----
/// Default tracing level (overridden by `RUST_LOG`).
#[arg(long, default_value = "info")]
pub log_level: String,
/// Log output format.
#[arg(long, value_enum, default_value = "text")]
pub log_format: LogFormat,
}
impl Cli {
/// Resolve parsed flags into a validated [`Config`].
///
/// Builds the [`DiscoveryBackend`] (enforcing static-vs-k8s mutual
/// exclusivity and resolving the k8s selector grammar via
/// [`resolve_mode`]), assembles the single [`ModelConfig`], then runs
/// [`Config::validate`] for the remaining value-level invariants
/// (model id, static worker URLs).
pub fn into_config(self) -> Result<Config> {
let discovery = self.build_discovery()?;
// Reject knobs that only take effect alongside another flag, rather
// than silently dropping them — mirrors the discovery mutual-exclusion
// checks. Otherwise an operator believes they tuned something that has
// no effect.
if self.cb_cool_down_secs.is_some() && self.cb_threshold.is_none() {
return Err(anyhow!(
"--cb-cool-down-secs requires --cb-threshold (the circuit breaker is \
enabled by --cb-threshold)"
));
}
let tuned_cache_aware = self.cache_threshold.is_some()
|| self.balance_abs_threshold.is_some()
|| self.balance_rel_threshold.is_some();
if tuned_cache_aware && self.policy != PolicyKind::CacheAwareZmq {
return Err(anyhow!(
"--cache-threshold / --balance-abs-threshold / --balance-rel-threshold \
require --policy cache_aware_zmq"
));
}
let circuit_breaker = self.cb_threshold.map(|threshold| CircuitBreakerConfig {
threshold,
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.
let cache_aware = if tuned_cache_aware {
let d = CacheAwareConfig::default();
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),
})
} else {
None
};
let config = Config {
server: ServerConfig {
host: self.host,
port: self.port,
},
observability: ObservabilityConfig {
log_level: self.log_level,
log_format: self.log_format,
},
model: ModelConfig {
// Default the tokenizer source to the model id (treated as a
// HuggingFace repo id) when --tokenizer-path is omitted.
tokenizer_path: self.tokenizer_path.unwrap_or_else(|| self.model_id.clone()),
id: self.model_id,
policy: self.policy,
circuit_breaker,
cache_aware,
},
discovery,
proxy: ProxyConfig {
request_timeout_secs: self.request_timeout_secs,
},
active_load: ActiveLoadConfig {
stale_request_timeout_secs: self.stale_request_timeout_secs,
},
};
config.validate()?;
Ok(config)
}
/// Resolve the discovery flags into a [`DiscoveryBackend`].
///
/// `--worker-urls` (static) and `--service-discovery` (k8s) are
/// mutually exclusive and exactly one is required. K8s-only flags
/// passed without `--service-discovery` are rejected so a typo can't
/// silently fall back to the static (empty) path. The k8s selector
/// grammar (plain vs PD) is validated eagerly here by [`resolve_mode`]
/// before the `K8sDiscoveryConfig` is constructed, so an invalid
/// combination is never stored.
fn build_discovery(&self) -> Result<DiscoveryBackend> {
let has_static = !self.worker_urls.is_empty();
let backend = match (has_static, self.service_discovery) {
(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()
|| !self.selector.is_empty()
|| !self.prefill_selector.is_empty()
|| !self.decode_selector.is_empty()
{
return Err(anyhow!(
"--service-discovery-namespace / --selector / --prefill-selector / \
--decode-selector require --service-discovery"
));
}
DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: self.worker_urls.clone(),
})
}
(false, true) => {
// Resolve (and validate) the selector flags into a
// K8sDiscoveryMode here, so an invalid combination can't be
// stored. Surfaces ConfigError as anyhow for the CLI.
let mode = resolve_mode(
join_selector(&self.selector).as_deref(),
join_selector(&self.prefill_selector).as_deref(),
join_selector(&self.decode_selector).as_deref(),
)
.map_err(|e| anyhow!("{e}"))?;
DiscoveryBackend::K8s(K8sDiscoveryConfig {
namespace: self.service_discovery_namespace.clone().unwrap_or_default(),
mode,
})
}
};
Ok(backend)
}
}
/// 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
/// its plain-vs-PD rules (and surface `NoSelector`).
fn join_selector(terms: &[String]) -> Option<String> {
if terms.is_empty() {
None
} else {
Some(terms.join(","))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{DiscoveryBackend, K8sDiscoveryMode};
/// Parse argv (without the leading binary name) into a `Config`.
fn into_config(args: &[&str]) -> Result<Config> {
let argv = std::iter::once("sgl-router").chain(args.iter().copied());
let cli = Cli::try_parse_from(argv).map_err(|e| anyhow!("{e}"))?;
cli.into_config()
}
const MODEL_ARGS: &[&str] = &[
"--model-id",
"qwen3-0.6b",
"--tokenizer-path",
"/tmp/qwen.json",
];
fn with_model(extra: &[&str]) -> Vec<String> {
MODEL_ARGS
.iter()
.chain(extra.iter())
.map(|s| s.to_string())
.collect()
}
fn into_config_owned(args: Vec<String>) -> Result<Config> {
let refs: Vec<&str> = args.iter().map(String::as_str).collect();
into_config(&refs)
}
#[test]
fn defaults_host_port_and_policy() {
let c = into_config_owned(with_model(&["--worker-urls", "http://10.0.0.1:30000"])).unwrap();
assert_eq!(c.server.host, "127.0.0.1");
assert_eq!(c.server.port, 30000);
assert_eq!(c.model.policy, PolicyKind::RoundRobin);
assert_eq!(c.model.id, "qwen3-0.6b");
assert_eq!(c.proxy.request_timeout_secs, 300);
assert_eq!(c.active_load.stale_request_timeout_secs, 600);
}
/// With `--tokenizer-path` omitted, the tokenizer source defaults to the
/// model id (treated as an HF repo id at load time).
#[test]
fn tokenizer_path_defaults_to_model_id_when_omitted() {
let c = into_config(&[
"--model-id",
"Qwen/Qwen3-0.6B",
"--worker-urls",
"http://x:30000",
])
.unwrap();
assert_eq!(c.model.id, "Qwen/Qwen3-0.6B");
assert_eq!(c.model.tokenizer_path, "Qwen/Qwen3-0.6B");
}
#[test]
fn explicit_tokenizer_path_is_used() {
let c = into_config(&[
"--model-id",
"qwen3",
"--tokenizer-path",
"/models/qwen3/tokenizer.json",
"--worker-urls",
"http://x:30000",
])
.unwrap();
assert_eq!(c.model.tokenizer_path, "/models/qwen3/tokenizer.json");
}
#[test]
fn static_urls_backend() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://10.0.0.1:30000",
"http://10.0.0.2:30000",
]))
.unwrap();
match &c.discovery {
DiscoveryBackend::StaticUrls(s) => assert_eq!(
s.urls,
vec![
"http://10.0.0.1:30000".to_string(),
"http://10.0.0.2:30000".to_string()
]
),
_ => panic!("expected static_urls backend"),
}
}
#[test]
fn rejects_no_discovery_backend() {
let err = into_config_owned(with_model(&[])).unwrap_err().to_string();
assert!(err.contains("no discovery backend"), "got: {err}");
}
#[test]
fn rejects_both_discovery_backends() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--service-discovery",
]))
.unwrap_err()
.to_string();
assert!(err.contains("mutually exclusive"), "got: {err}");
}
#[test]
fn rejects_k8s_flags_without_service_discovery() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--selector",
"app=sglang",
]))
.unwrap_err()
.to_string();
assert!(err.contains("require --service-discovery"), "got: {err}");
}
#[test]
fn rejects_static_urls_duplicate() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"http://x:30000",
]))
.unwrap_err()
.to_string();
assert!(err.contains("duplicate"), "got: {err}");
}
#[test]
fn rejects_static_urls_schemeless() {
let err = into_config_owned(with_model(&["--worker-urls", "10.0.0.1:30000"]))
.unwrap_err()
.to_string();
assert!(
err.contains("not a valid URL") || err.contains("unsupported scheme"),
"got: {err}"
);
}
#[test]
fn rejects_static_urls_non_http_scheme() {
let err = into_config_owned(with_model(&["--worker-urls", "ws://x:30000"]))
.unwrap_err()
.to_string();
assert!(err.contains("unsupported scheme"), "got: {err}");
}
#[test]
fn k8s_plain_backend() {
let c = into_config_owned(with_model(&[
"--service-discovery",
"--service-discovery-namespace",
"prod",
"--selector",
"app=engines-qwen3",
]))
.unwrap();
match &c.discovery {
DiscoveryBackend::K8s(k) => {
assert_eq!(k.namespace, "prod");
assert_eq!(
k.mode,
K8sDiscoveryMode::Plain {
label_selector: "app=engines-qwen3".to_string()
}
);
}
_ => panic!("expected k8s backend"),
}
}
/// Multiple `--selector` terms AND-join into one comma-separated
/// label selector (matches the Python router's space-separated form).
#[test]
fn k8s_plain_selector_joins_multiple_terms() {
let c = into_config_owned(with_model(&[
"--service-discovery",
"--selector",
"app=sglang",
"zone=us-east",
]))
.unwrap();
match &c.discovery {
DiscoveryBackend::K8s(k) => assert_eq!(
k.mode,
K8sDiscoveryMode::Plain {
label_selector: "app=sglang,zone=us-east".to_string()
}
),
_ => panic!("expected k8s backend"),
}
}
/// Empty namespace is intentional — it triggers a cluster-wide watch.
#[test]
fn k8s_empty_namespace_watches_all() {
let c = into_config_owned(with_model(&[
"--service-discovery",
"--selector",
"app=sglang",
]))
.unwrap();
match &c.discovery {
DiscoveryBackend::K8s(k) => assert_eq!(k.namespace, ""),
_ => panic!("expected k8s backend"),
}
}
#[test]
fn k8s_pd_backend() {
let c = into_config_owned(with_model(&[
"--service-discovery",
"--service-discovery-namespace",
"default",
"--prefill-selector",
"app=sglang,role=prefill",
"--decode-selector",
"app=sglang,role=decode",
]))
.unwrap();
match &c.discovery {
DiscoveryBackend::K8s(k) => assert_eq!(
k.mode,
K8sDiscoveryMode::PdDisaggregation {
prefill_selector: "app=sglang,role=prefill".to_string(),
decode_selector: "app=sglang,role=decode".to_string(),
}
),
_ => panic!("expected k8s backend"),
}
}
/// `--service-discovery` with no selector at all fails `resolve_mode`
/// validation with the `NoSelector` wording.
#[test]
fn rejects_k8s_without_selector() {
let err = into_config_owned(with_model(&["--service-discovery"]))
.unwrap_err()
.to_string()
.to_lowercase();
assert!(err.contains("none were set"), "got: {err}");
}
/// `--prefill-selector` without `--decode-selector` is rejected through
/// the full CLI path — pins that `build_discovery` feeds the right
/// selectors into `resolve_mode` (a positional mix-up would surface a
/// different error or none).
#[test]
fn rejects_k8s_partial_pd_selectors() {
let err = into_config_owned(with_model(&[
"--service-discovery",
"--prefill-selector",
"app=sglang,role=prefill",
]))
.unwrap_err()
.to_string();
assert!(
err.contains("PD mode requires BOTH"),
"expected PartialPdSelectors wording, got: {err}"
);
}
/// Identical prefill/decode selectors are rejected through the full CLI
/// path (would silently leave the decode pool empty at runtime).
#[test]
fn rejects_k8s_identical_pd_selectors() {
let err = into_config_owned(with_model(&[
"--service-discovery",
"--prefill-selector",
"app=sglang",
"--decode-selector",
"app=sglang",
]))
.unwrap_err()
.to_string();
assert!(
err.contains("must differ"),
"expected IdenticalPdSelectors wording, got: {err}"
);
}
/// clap rejects an unknown `--policy` value at parse time.
#[test]
fn rejects_unknown_policy() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--policy",
"bogus_policy",
]))
.unwrap_err()
.to_string();
assert!(
err.contains("bogus_policy") || err.contains("policy"),
"got: {err}"
);
}
/// clap rejects `--cb-threshold 0` because the field is `NonZeroU32`.
#[test]
fn rejects_zero_cb_threshold() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--cb-threshold",
"0",
]))
.unwrap_err()
.to_string();
assert!(err.contains("cb-threshold"), "got: {err}");
}
#[test]
fn cb_threshold_enables_circuit_breaker_with_default_cool_down() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--cb-threshold",
"5",
]))
.unwrap();
let cb = c.model.circuit_breaker.expect("cb enabled");
assert_eq!(cb.threshold.get(), 5);
assert_eq!(cb.cool_down_secs, 30);
}
#[test]
fn cb_cool_down_honors_explicit_override() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--cb-threshold",
"3",
"--cb-cool-down-secs",
"10",
]))
.unwrap();
let cb = c.model.circuit_breaker.expect("cb enabled");
assert_eq!(cb.cool_down_secs, 10);
}
#[test]
fn rejects_cb_cool_down_without_threshold() {
let err = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--cb-cool-down-secs",
"10",
]))
.unwrap_err()
.to_string();
assert!(
err.contains("--cb-cool-down-secs requires --cb-threshold"),
"got: {err}"
);
}
#[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 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(&[
"--worker-urls",
"http://x:30000",
"--log-format",
"json",
]))
.unwrap();
assert_eq!(c.observability.log_format, LogFormat::Json);
}
/// Pins that the two timeout overrides land in the right fields — they
/// are adjacent `u64`s with similar names, so a copy-paste swap would
/// otherwise go unnoticed (and `stale` must sit above `proxy`).
#[test]
fn timeout_overrides_land_in_distinct_fields() {
let c = into_config_owned(with_model(&[
"--worker-urls",
"http://x:30000",
"--request-timeout-secs",
"120",
"--stale-request-timeout-secs",
"240",
]))
.unwrap();
assert_eq!(c.proxy.request_timeout_secs, 120);
assert_eq!(c.active_load.stale_request_timeout_secs, 240);
}
}
+67 -453
View File
@@ -1,42 +1,21 @@
pub mod cli;
pub mod types; pub mod types;
pub use cli::Cli;
pub use types::*; pub use types::*;
use anyhow::Context as _;
use anyhow::{anyhow, Result}; use anyhow::{anyhow, Result};
use std::path::Path;
impl Config { impl Config {
pub fn from_path(p: &Path) -> Result<Self> { /// Check invariants the type system and `clap` don't already enforce.
let raw = /// Called by [`cli::Cli::into_config`] after assembling the `Config`
std::fs::read_to_string(p).with_context(|| format!("read config {}", p.display()))?; /// from flags. Unknown policy names and `--cb-threshold 0` are
let ext = p.extension().and_then(|s| s.to_str()).unwrap_or(""); /// rejected at parse time (`ValueEnum` / `NonZeroU32`); only the
let cfg: Config = match ext { /// remaining value-level invariants are checked here.
"yaml" | "yml" => serde_yaml::from_str(&raw) pub(crate) fn validate(&self) -> Result<()> {
.map_err(|e| anyhow!("parse yaml {}: {e}", p.display()))?, if self.model.id.is_empty() {
"toml" => { return Err(anyhow!("model id must be non-empty"));
toml::from_str(&raw).map_err(|e| anyhow!("parse toml {}: {e}", p.display()))?
} }
other => { match &self.discovery {
return Err(anyhow!(
"unsupported config extension {other:?}; want yaml/yml/toml"
))
}
};
cfg.validate()?;
Ok(cfg)
}
fn validate(&self) -> Result<()> {
// Unknown policy names are rejected by serde via `PolicyKind`'s
// `rename_all = "snake_case"`; threshold = 0 is rejected by
// `NonZeroU32`. Only fields without a type-system constraint are
// checked here.
for m in &self.models {
if m.id.is_empty() {
return Err(anyhow!("model.id must be non-empty"));
}
}
match &self.discovery.backend {
DiscoveryBackend::StaticUrls(s) => { DiscoveryBackend::StaticUrls(s) => {
if s.urls.is_empty() { if s.urls.is_empty() {
return Err(anyhow!( return Err(anyhow!(
@@ -44,7 +23,7 @@ impl Config {
)); ));
} }
// Validate every entry up front so typos surface at // Validate every entry up front so typos surface at
// config-load with a precise diagnostic instead of as // startup with a precise diagnostic instead of as
// per-worker introspect failures or as two registry // per-worker introspect failures or as two registry
// entries pointing at the same SGLang (trailing-slash // entries pointing at the same SGLang (trailing-slash
// near-duplicates). Dedupe runs against a normalized // near-duplicates). Dedupe runs against a normalized
@@ -77,14 +56,11 @@ impl Config {
} }
} }
} }
DiscoveryBackend::K8s(k) => { // K8s selector validity is resolved at construction time
// Empty namespace is intentional: triggers `Api::all(client)` // (`resolve_mode` in `Cli::build_discovery`), so the stored
// for cluster-wide EndpointSlice watch (see // `K8sDiscoveryMode` is already valid here. Any namespace
// `discovery::k8s::spawn`). Only validate the selector // (including empty, for a cluster-wide watch) is accepted.
// combination here. DiscoveryBackend::K8s(_) => {}
let _ = &k.namespace;
k.mode().map_err(|e| anyhow!("{e}"))?;
}
} }
Ok(()) Ok(())
} }
@@ -94,447 +70,85 @@ impl Config {
mod tests { mod tests {
use super::*; use super::*;
/// Write `body` to a temp file with the given extension and load it /// Build a minimal valid-shape `Config` with the given static worker
/// through `Config::from_path`. Failures still surface the offending /// URLs and model id, so the `validate()` branches can be exercised
/// config because each call site passes its body inline. /// directly. CLI parsing and the static-vs-k8s mapping are covered in
fn load(ext: &str, body: &str) -> Result<Config> { /// the `cli` module tests; the k8s selector grammar in `types`.
let dir = tempfile::tempdir().unwrap(); fn cfg(model_id: &str, urls: &[&str]) -> Config {
let p = dir.path().join(format!("c.{ext}")); Config {
std::fs::write(&p, body).unwrap(); server: ServerConfig {
Config::from_path(&p) host: "127.0.0.1".into(),
} port: 30000,
},
#[test] observability: ObservabilityConfig::default(),
fn loads_minimal_yaml() { model: ModelConfig {
let c = load( id: model_id.into(),
"yaml", tokenizer_path: "/tmp/tok.json".into(),
r#" policy: PolicyKind::RoundRobin,
server: circuit_breaker: None,
host: "0.0.0.0" cache_aware: None,
port: 8090 },
models: discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
- id: "qwen3-0.6b" urls: urls.iter().map(|s| s.to_string()).collect(),
tokenizer_path: "/tmp/qwen.json" }),
discovery: proxy: ProxyConfig::default(),
backend: static_urls active_load: ActiveLoadConfig::default(),
static_urls:
urls:
- "http://10.0.0.1:30000"
"#,
)
.unwrap();
assert_eq!(c.server.port, 8090);
assert_eq!(c.models[0].id, "qwen3-0.6b");
match &c.discovery.backend {
DiscoveryBackend::StaticUrls(s) => {
assert_eq!(s.urls, vec!["http://10.0.0.1:30000".to_string()])
}
_ => panic!("expected static_urls backend"),
} }
} }
#[test] #[test]
fn loads_minimal_toml() { fn accepts_minimal_static_config() {
let c = load( cfg("qwen3", &["http://10.0.0.1:30000"]).validate().unwrap();
"toml",
r#"
[server]
host = "0.0.0.0"
port = 8090
[[models]]
id = "qwen3-0.6b"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://10.0.0.1:30000"]
"#,
)
.unwrap();
assert_eq!(c.server.port, 8090);
match &c.discovery.backend {
DiscoveryBackend::StaticUrls(s) => {
assert_eq!(s.urls, vec!["http://10.0.0.1:30000".to_string()])
}
_ => panic!("expected static_urls backend"),
}
} }
#[test] #[test]
fn rejects_missing_discovery_section() { fn rejects_empty_model_id() {
let err = load( let err = cfg("", &["http://10.0.0.1:30000"])
"yaml", .validate()
"server:\n host: \"0.0.0.0\"\n port: 8090\nmodels: []\n",
)
.unwrap_err();
let msg = err.to_string().to_lowercase();
assert!(
msg.contains("discovery") || msg.contains("missing"),
"got: {err}"
);
}
#[test]
fn rejects_unknown_extension() {
let err = load("txt", "").unwrap_err();
assert!(err.to_string().contains("yaml") && err.to_string().contains("toml"));
}
#[test]
fn loads_static_urls_discovery() {
let c = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "qwen3-0.6b"
tokenizer_path = "/tmp/qwen.json"
policy = "round_robin"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://10.0.0.1:30000", "http://10.0.0.2:30000"]
"#,
)
.unwrap();
match &c.discovery.backend {
DiscoveryBackend::StaticUrls(s) => {
assert_eq!(
s.urls,
vec![
"http://10.0.0.1:30000".to_string(),
"http://10.0.0.2:30000".to_string(),
],
);
}
_ => panic!("expected static_urls backend"),
}
assert_eq!(c.models[0].policy, PolicyKind::RoundRobin);
}
#[test]
fn rejects_static_urls_with_empty_list() {
let err = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = []
"#,
)
.unwrap_err() .unwrap_err()
.to_string(); .to_string();
assert!(err.contains("model id"), "got: {err}");
}
#[test]
fn rejects_empty_static_urls_list() {
let err = cfg("qwen3", &[]).validate().unwrap_err().to_string();
assert!(err.contains("non-empty"), "got: {err}"); assert!(err.contains("non-empty"), "got: {err}");
} }
#[test] #[test]
fn rejects_static_urls_with_duplicate_entry() { fn rejects_static_urls_empty_entry() {
let err = load( let err = cfg("qwen3", &["http://x:30000", ""])
"toml", .validate()
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://x:30000", "http://x:30000"]
"#,
)
.unwrap_err()
.to_string();
assert!(err.contains("duplicate"), "got: {err}");
}
#[test]
fn rejects_static_urls_with_empty_entry() {
let err = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://x:30000", ""]
"#,
)
.unwrap_err() .unwrap_err()
.to_string(); .to_string();
assert!(err.contains("empty"), "got: {err}"); assert!(err.contains("empty"), "got: {err}");
} }
/// Whitespace-only entries are user typos that previously slipped
/// through `is_empty()` checks and surfaced as "introspect against
/// ` /server_info` failed" at runtime. Catch at load.
#[test] #[test]
fn rejects_static_urls_with_whitespace_only_entry() { fn rejects_static_urls_whitespace_only_entry() {
let err = load( let err = cfg("qwen3", &["http://x:30000", " "])
"toml", .validate()
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://x:30000", " "]
"#,
)
.unwrap_err() .unwrap_err()
.to_string(); .to_string();
assert!(err.contains("whitespace"), "got: {err}"); assert!(err.contains("empty or whitespace"), "got: {err}");
} }
/// `"10.0.0.1:30000"` (missing scheme) used to pass validation; the
/// scheme/`http://` would only fail (or worse, silently degrade
/// because of the `parse_bootstrap_host` localhost fallback) at
/// introspect time. Reject at load.
#[test] #[test]
fn rejects_static_urls_with_schemeless_entry() { fn rejects_static_urls_trailing_slash_near_duplicate() {
let err = load( let err = cfg("qwen3", &["http://x:30000", "http://x:30000/"])
"toml", .validate()
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["10.0.0.1:30000"]
"#,
)
.unwrap_err()
.to_string();
assert!(
err.contains("not a valid URL") || err.contains("unsupported scheme"),
"got: {err}"
);
}
/// Non-http(s) schemes are rejected. The router speaks HTTP to
/// workers; a `tcp://` or `ws://` entry is almost certainly an
/// operator typo.
#[test]
fn rejects_static_urls_with_non_http_scheme() {
let err = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["ws://x:30000"]
"#,
)
.unwrap_err()
.to_string();
assert!(err.contains("unsupported scheme"), "got: {err}");
}
/// Trailing-slash near-duplicates collide in the registry but used
/// to pass byte-equality dedupe. Normalize before checking so two
/// pointers at the same SGLang surface as a config error.
#[test]
fn rejects_static_urls_with_trailing_slash_near_duplicate() {
let err = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "m"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://x:30000", "http://x:30000/"]
"#,
)
.unwrap_err() .unwrap_err()
.to_string(); .to_string();
assert!(err.contains("duplicate"), "got: {err}"); assert!(err.contains("duplicate"), "got: {err}");
} }
#[test] #[test]
fn loads_k8s_discovery() { fn rejects_static_urls_non_http_scheme() {
let c = load( let err = cfg("qwen3", &["ws://x:30000"])
"toml", .validate()
r#" .unwrap_err()
[server] .to_string();
host = "127.0.0.1" assert!(err.contains("unsupported scheme"), "got: {err}");
port = 8090
[[models]]
id = "qwen3-0.6b"
tokenizer_path = "/tmp/qwen.json"
policy = "round_robin"
[discovery]
backend = "k8s"
[discovery.k8s]
namespace = "default"
label_selector = "app=sglang"
"#,
)
.unwrap();
match &c.discovery.backend {
DiscoveryBackend::K8s(k) => {
assert_eq!(k.namespace, "default");
assert_eq!(k.label_selector.as_deref(), Some("app=sglang"));
assert!(k.prefill_selector.is_none());
assert!(k.decode_selector.is_none());
}
_ => panic!("expected k8s backend"),
}
}
/// K8s PD selectors drive slice-classification only; per-worker
/// bootstrap_port comes from `/server_info` post-discovery
/// (`crate::workers::introspect`). This test pins the wire-shape;
/// the selector grammar itself is covered in `types.rs`.
#[test]
fn loads_k8s_pd_discovery_with_prefill_and_decode_selectors() {
let c = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "qwen3-0.6b"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "k8s"
[discovery.k8s]
namespace = "default"
prefill_selector = "app=sglang,role=prefill"
decode_selector = "app=sglang,role=decode"
"#,
)
.expect("k8s PD config must load");
match &c.discovery.backend {
DiscoveryBackend::K8s(k) => {
assert_eq!(k.namespace, "default");
assert_eq!(
k.prefill_selector.as_deref(),
Some("app=sglang,role=prefill")
);
assert_eq!(k.decode_selector.as_deref(), Some("app=sglang,role=decode"));
assert!(k.label_selector.is_none());
}
_ => panic!("expected k8s backend"),
}
}
#[test]
fn rejects_k8s_config_with_no_selector() {
let err = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "qwen"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "k8s"
[discovery.k8s]
namespace = "default"
"#,
)
.unwrap_err();
// Pin the specific variant: `ConfigError::NoSelector` ("none were
// set"). A bare `contains("selector")` would also pass for
// EmptyPdSelector / PartialPdSelectors / IdenticalPdSelectors /
// UnsupportedSelectorGrammar — variants that have semantically
// different error wording but all mention "selector". A future
// regression that returned, say, `PartialPdSelectors` for the
// all-None input would be caught here.
let msg = err.to_string().to_lowercase();
assert!(
msg.contains("none were set"),
"expected NoSelector wording (\"none were set\"); got: {err}",
);
}
// Direct `K8sDiscoveryConfig::mode()` unit tests live alongside the
// type in `src/config/types.rs::k8s_discovery_config_tests`.
// The tests in this module exercise the `Config::from_path` ↔ K8s
// selector wiring, not the selector grammar itself.
#[test]
fn rejects_unknown_policy_name() {
let err = load(
"yaml",
"
server:
host: 0.0.0.0
port: 8090
discovery:
backend: static_urls
static_urls:
urls:
- http://x:30000
models:
- id: qwen
tokenizer_path: /tmp/qwen.json
policy: bogus_policy
",
)
.unwrap_err();
let msg = err.to_string().to_lowercase();
assert!(
msg.contains("bogus_policy") || msg.contains("policy"),
"got: {err}"
);
}
#[test]
fn defaults_policy_to_round_robin() {
let c = load(
"toml",
r#"
[server]
host = "127.0.0.1"
port = 8090
[[models]]
id = "qwen"
tokenizer_path = "/tmp/qwen.json"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["http://x:30000"]
"#,
)
.unwrap();
assert_eq!(c.models[0].policy, PolicyKind::RoundRobin);
} }
} }
+111 -255
View File
@@ -1,32 +1,35 @@
use serde::{Deserialize, Serialize};
use std::num::NonZeroU32; use std::num::NonZeroU32;
#[derive(Debug, Clone, Serialize, Deserialize)] /// In-memory router configuration, built from CLI flags by
/// [`crate::config::cli::Cli::into_config`] and validated by
/// [`Config::validate`]. The router serves exactly one model.
#[derive(Debug, Clone)]
pub struct Config { pub struct Config {
pub server: ServerConfig, pub server: ServerConfig,
#[serde(default)]
pub observability: ObservabilityConfig, pub observability: ObservabilityConfig,
pub models: Vec<ModelConfig>, pub model: ModelConfig,
pub discovery: DiscoveryConfig, /// Selected discovery backend. Built from CLI flags by
#[serde(default)] /// [`crate::config::cli::Cli::into_config`]: the static-vs-k8s choice
/// and the k8s selector grammar are resolved there (the latter via
/// [`resolve_mode`]); static worker-URL validity is checked by
/// [`Config::validate`].
pub discovery: DiscoveryBackend,
pub proxy: ProxyConfig, pub proxy: ProxyConfig,
#[serde(default)]
pub active_load: ActiveLoadConfig, pub active_load: ActiveLoadConfig,
} }
/// Outbound proxy tuning. Default mirrors SGLang's typical prefill / /// Outbound proxy tuning. Default mirrors SGLang's typical prefill /
/// decode latency budget; e2e tests lower it so per-request failures /// decode latency budget; e2e tests lower it so per-request failures
/// trip the circuit breaker within the test's wall-time. /// trip the circuit breaker within the test's wall-time.
#[derive(Debug, Clone, Copy, Serialize, Deserialize)] #[derive(Debug, Clone, Copy)]
pub struct ProxyConfig { pub struct ProxyConfig {
/// Maximum time to wait for a single upstream HTTP request to /// Maximum time to wait for a single upstream HTTP request to
/// return headers + body. Default 300 s. The circuit breaker /// return headers + body. Default 300 s. The circuit breaker
/// records a failure when this fires. /// records a failure when this fires.
#[serde(default = "default_proxy_request_timeout_secs")]
pub request_timeout_secs: u64, pub request_timeout_secs: u64,
} }
fn default_proxy_request_timeout_secs() -> u64 { pub fn default_proxy_request_timeout_secs() -> u64 {
300 300
} }
@@ -42,16 +45,15 @@ impl Default for ProxyConfig {
/// sits above `proxy.request_timeout_secs` so the proxy timeout is the /// sits above `proxy.request_timeout_secs` so the proxy timeout is the
/// one users hit first for normal slow upstreams; tests lower it to /// one users hit first for normal slow upstreams; tests lower it to
/// let the janitor fire within their wall-time budget. /// let the janitor fire within their wall-time budget.
#[derive(Debug, Clone, Copy, Serialize, Deserialize)] #[derive(Debug, Clone, Copy)]
pub struct ActiveLoadConfig { pub struct ActiveLoadConfig {
/// How long a request entry can live in the registry before the /// How long a request entry can live in the registry before the
/// janitor fires its `cancel_token` and the chat handler returns /// janitor fires its `cancel_token` and the chat handler returns
/// 504 `stale_request_expired`. Default 600 s. /// 504 `stale_request_expired`. Default 600 s.
#[serde(default = "default_stale_request_timeout_secs")]
pub stale_request_timeout_secs: u64, pub stale_request_timeout_secs: u64,
} }
fn default_stale_request_timeout_secs() -> u64 { pub fn default_stale_request_timeout_secs() -> u64 {
600 600
} }
@@ -63,51 +65,52 @@ impl Default for ActiveLoadConfig {
} }
} }
/// Routing policy selector — the enum form lets serde reject unknown /// Routing policy selector — the enum form lets `clap` reject unknown
/// values at deserialization time and removes the runtime string match in /// values at parse time and removes the runtime string match in the
/// the policy factory. /// policy factory.
/// ///
/// Serialised as `"round_robin"` / `"random"` / `"power_of_two"` / /// Accepted on the CLI (`--policy`) as `round_robin` / `random` /
/// `"cache_aware_zmq"`. /// `power_of_two` / `cache_aware_zmq`.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Default, clap::ValueEnum)]
#[serde(rename_all = "snake_case")]
pub enum PolicyKind { pub enum PolicyKind {
#[default] #[default]
#[value(name = "round_robin")]
RoundRobin, RoundRobin,
#[value(name = "random")]
Random, Random,
#[value(name = "power_of_two")]
PowerOfTwo, PowerOfTwo,
/// Cache-aware routing fed by SGLang's ZMQ KV-cache event publisher. /// Cache-aware routing fed by SGLang's ZMQ KV-cache event publisher.
/// Requires the model to have a tokenizer loaded; cache_aware tuning /// Requires the model to have a tokenizer loaded; cache_aware tuning
/// lives on `ModelConfig::cache_aware`. /// lives on `ModelConfig::cache_aware`.
#[value(name = "cache_aware_zmq")]
CacheAwareZmq, CacheAwareZmq,
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone)]
pub struct ServerConfig { pub struct ServerConfig {
pub host: String, pub host: String,
pub port: u16, pub port: u16,
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone)]
pub struct ObservabilityConfig { pub struct ObservabilityConfig {
#[serde(default = "default_log_level")]
pub log_level: String, pub log_level: String,
/// Selects the tracing-subscriber output format. Serde rejects /// Selects the tracing-subscriber output format. `clap` rejects
/// unrecognized values at config-load (`"jsonl"` and similar /// unrecognized values at parse time (`--log-format jsonl` and
/// plausible typos surface as an error instead of silently /// similar typos surface as an error instead of silently degrading
/// degrading to text), matching the discoverability pattern used /// to text).
/// by `policy` and `discovery.backend`.
#[serde(default)]
pub log_format: LogFormat, pub log_format: LogFormat,
} }
/// `text` for human-readable dev output, `json` for one-line-per-record /// `text` for human-readable dev output, `json` for one-line-per-record
/// JSON suitable for k8s log aggregators (fluent-bit / vector / Loki). /// JSON suitable for k8s log aggregators (fluent-bit / vector / Loki).
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, clap::ValueEnum)]
#[serde(rename_all = "lowercase")]
pub enum LogFormat { pub enum LogFormat {
#[default] #[default]
#[value(name = "text")]
Text, Text,
#[value(name = "json")]
Json, Json,
} }
@@ -124,40 +127,37 @@ impl Default for ObservabilityConfig {
} }
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone)]
pub struct ModelConfig { pub struct ModelConfig {
pub id: String, pub id: String,
/// Tokenizer source: a local `tokenizer.json` path or a HuggingFace repo
/// id (downloaded on demand). Defaults to `id` when `--tokenizer-path`
/// is omitted. Resolved by [`crate::tokenizer::adapter::load`].
pub tokenizer_path: String, pub tokenizer_path: String,
#[serde(default)]
pub policy: PolicyKind, pub policy: PolicyKind,
#[serde(default)]
pub circuit_breaker: Option<CircuitBreakerConfig>, pub circuit_breaker: Option<CircuitBreakerConfig>,
/// Tuning for the cache-aware ZMQ policy. Ignored unless /// Tuning for the cache-aware ZMQ policy. Ignored unless
/// `policy = "cache_aware_zmq"`. `None` falls back to defaults at /// `policy = "cache_aware_zmq"`. `None` falls back to defaults at
/// policy construction time. /// policy construction time.
#[serde(default)]
pub cache_aware: Option<CacheAwareConfig>, pub cache_aware: Option<CacheAwareConfig>,
} }
/// Per-model cache-aware-ZMQ tuning. /// Per-model cache-aware-ZMQ tuning.
#[derive(Debug, Clone, Copy, Serialize, Deserialize)] #[derive(Debug, Clone, Copy)]
pub struct CacheAwareConfig { pub struct CacheAwareConfig {
/// Lower bound on `matched_blocks / total_blocks` for the tree match /// Lower bound on `matched_blocks / total_blocks` for the tree match
/// to win the selection. Below this, the policy falls back to /// to win the selection. Below this, the policy falls back to
/// min-load. Default 0.5 — a half-cached prompt is still a strong /// min-load. Default 0.5 — a half-cached prompt is still a strong
/// signal but not so weak that random hash collisions could trigger /// signal but not so weak that random hash collisions could trigger
/// affinity to an arbitrary worker. /// affinity to an arbitrary worker.
#[serde(default = "default_cache_threshold")]
pub cache_threshold: f32, pub cache_threshold: f32,
/// Absolute load spread (`max - min`) above which the cache check is /// Absolute load spread (`max - min`) above which the cache check is
/// skipped in favour of min-load. Default 32 — picked to dominate /// skipped in favour of min-load. Default 32 — picked to dominate
/// over typical batch-of-8 effect. /// over typical batch-of-8 effect.
#[serde(default = "default_balance_abs")]
pub balance_abs_threshold: usize, pub balance_abs_threshold: usize,
/// Multiplicative load spread (`max > min * balance_rel_threshold`) /// Multiplicative load spread (`max > min * balance_rel_threshold`)
/// that the absolute check is gated on. Default 1.1 — 10 % relative /// that the absolute check is gated on. Default 1.1 — 10 % relative
/// difference triggers re-balancing. /// difference triggers re-balancing.
#[serde(default = "default_balance_rel")]
pub balance_rel_threshold: f32, pub balance_rel_threshold: f32,
} }
@@ -181,125 +181,22 @@ fn default_balance_rel() -> f32 {
1.1 1.1
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone)]
pub struct CircuitBreakerConfig { pub struct CircuitBreakerConfig {
/// Consecutive failures required before the breaker opens. Encoded /// Consecutive failures required before the breaker opens. Encoded
/// as `NonZeroU32` so a config setting `threshold = 0` (which would /// as `NonZeroU32` so `--cb-threshold 0` (which would open the
/// open the breaker before any failure) is rejected at deserialization /// breaker before any failure) is rejected at CLI-parse time rather
/// rather than silently behaving as "always open". /// than silently behaving as "always open".
#[serde(default = "default_cb_threshold")]
pub threshold: NonZeroU32, pub threshold: NonZeroU32,
#[serde(default = "default_cb_cool_down")]
pub cool_down_secs: u64, pub cool_down_secs: u64,
} }
fn default_cb_threshold() -> NonZeroU32 { /// Default circuit-breaker cool-down, applied when `--cb-threshold` is
NonZeroU32::new(3).unwrap() /// set without an explicit `--cb-cool-down-secs`.
} pub fn default_cb_cool_down() -> u64 {
fn default_cb_cool_down() -> u64 {
30 30
} }
/// Config-level discovery section. Deserialized from:
///
/// TOML:
/// ```toml
/// [discovery]
/// backend = "static_urls"
/// [discovery.static_urls]
/// urls = ["http://10.0.0.1:30000", "http://10.0.0.2:30000"]
/// ```
///
/// YAML:
/// ```yaml
/// discovery:
/// backend: static_urls
/// static_urls:
/// urls:
/// - http://10.0.0.1:30000
/// - http://10.0.0.2:30000
/// ```
///
/// The custom `Deserialize` impl on [`DiscoveryConfig`] converts the
/// raw fields into the resolved `DiscoveryBackend` enum via `try_from`.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DiscoveryConfigRaw {
pub backend: String,
pub static_urls: Option<StaticUrlsDiscoveryConfig>,
pub k8s: Option<K8sDiscoveryConfig>,
}
/// Post-validation discovery config with a resolved `DiscoveryBackend` enum.
/// Constructed by `Config::from_path` after `validate()`.
#[derive(Debug, Clone)]
pub struct DiscoveryConfig {
pub backend: DiscoveryBackend,
}
impl<'de> Deserialize<'de> for DiscoveryConfig {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw = DiscoveryConfigRaw::deserialize(deserializer)?;
raw.try_into().map_err(serde::de::Error::custom)
}
}
impl Serialize for DiscoveryConfig {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
let raw: DiscoveryConfigRaw = self.clone().into();
raw.serialize(serializer)
}
}
impl TryFrom<DiscoveryConfigRaw> for DiscoveryConfig {
type Error = String;
fn try_from(raw: DiscoveryConfigRaw) -> Result<Self, Self::Error> {
let backend = match raw.backend.as_str() {
"static_urls" => {
let s = raw.static_urls.ok_or(
"discovery.backend = \"static_urls\" requires [discovery.static_urls] section",
)?;
DiscoveryBackend::StaticUrls(s)
}
"k8s" => {
let k = raw
.k8s
.ok_or("discovery.backend = \"k8s\" requires [discovery.k8s] section")?;
DiscoveryBackend::K8s(k)
}
other => {
return Err(format!(
"unknown discovery.backend = {other:?}; valid: \"static_urls\", \"k8s\""
))
}
};
Ok(DiscoveryConfig { backend })
}
}
impl From<DiscoveryConfig> for DiscoveryConfigRaw {
fn from(cfg: DiscoveryConfig) -> Self {
match cfg.backend {
DiscoveryBackend::StaticUrls(s) => DiscoveryConfigRaw {
backend: "static_urls".to_string(),
static_urls: Some(s),
k8s: None,
},
DiscoveryBackend::K8s(k) => DiscoveryConfigRaw {
backend: "k8s".to_string(),
static_urls: None,
k8s: Some(k),
},
}
}
}
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub enum DiscoveryBackend { pub enum DiscoveryBackend {
StaticUrls(StaticUrlsDiscoveryConfig), StaticUrls(StaticUrlsDiscoveryConfig),
@@ -311,30 +208,25 @@ pub enum DiscoveryBackend {
/// from `/server_info` (see [`crate::workers::introspect`]). /// from `/server_info` (see [`crate::workers::introspect`]).
/// ///
/// No file watcher, no hot-reload: topology change requires a restart. /// No file watcher, no hot-reload: topology change requires a restart.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone)]
pub struct StaticUrlsDiscoveryConfig { pub struct StaticUrlsDiscoveryConfig {
pub urls: Vec<String>, pub urls: Vec<String>,
} }
/// Configuration for the Kubernetes `EndpointSlice` discovery backend. /// Configuration for the Kubernetes `EndpointSlice` discovery backend.
/// Built from the `--service-discovery*` / `--selector` / `--prefill-selector`
/// / `--decode-selector` flags by [`crate::config::cli::Cli::build_discovery`].
/// ///
/// Two operating modes, distinguished by which selector fields are set: /// Two operating modes, distinguished by which selector flags are set:
/// ///
/// 1. **Plain** — all matched workers share the same role: /// 1. **Plain** — all matched workers share the same role:
/// ```toml /// `--service-discovery-namespace default --selector app=sglang`
/// [discovery.k8s]
/// namespace = "default"
/// label_selector = "app=sglang"
/// ```
/// ///
/// 2. **PD disaggregation** — prefill and decode workers are separated by /// 2. **PD disaggregation** — prefill and decode workers are separated by
/// different selectors: /// different selectors:
/// ```toml /// `--service-discovery-namespace default
/// [discovery.k8s] /// --prefill-selector app=sglang,role=prefill
/// namespace = "default" /// --decode-selector app=sglang,role=decode`
/// prefill_selector = "app=sglang,role=prefill"
/// decode_selector = "app=sglang,role=decode"
/// ```
/// ///
/// In PD mode, the selectors drive **slice-classification** (which /// In PD mode, the selectors drive **slice-classification** (which
/// EndpointSlices feed the prefill pool vs the decode pool). The actual /// EndpointSlices feed the prefill pool vs the decode pool). The actual
@@ -344,20 +236,19 @@ pub struct StaticUrlsDiscoveryConfig {
/// [`crate::workers::introspect`] for the `disaggregation_mode` and /// [`crate::workers::introspect`] for the `disaggregation_mode` and
/// `disaggregation_bootstrap_port` extraction. /// `disaggregation_bootstrap_port` extraction.
/// ///
/// `mode()` validates the combination and returns the resolved /// [`resolve_mode`] validates the selector flags and produces the
/// [`K8sDiscoveryMode`]; any other selector combination is rejected. /// resolved [`K8sDiscoveryMode`] once, at construction in
#[derive(Debug, Clone, Serialize, Deserialize)] /// [`crate::config::cli::Cli::build_discovery`] — so an invalid selector
/// combination is unrepresentable here.
#[derive(Debug, Clone)]
pub struct K8sDiscoveryConfig { pub struct K8sDiscoveryConfig {
pub namespace: String, pub namespace: String,
#[serde(default)] /// Resolved + validated selector mode (plain vs PD).
pub label_selector: Option<String>, pub mode: K8sDiscoveryMode,
#[serde(default)]
pub prefill_selector: Option<String>,
#[serde(default)]
pub decode_selector: Option<String>,
} }
/// Resolved discovery mode derived from a [`K8sDiscoveryConfig`]. /// Resolved discovery mode, produced by [`resolve_mode`] from the CLI
/// selector flags and stored on [`K8sDiscoveryConfig`].
/// ///
/// The discovery backend uses this to: /// The discovery backend uses this to:
/// * pick the server-side `LIST` label selector (Plain: the single selector; /// * pick the server-side `LIST` label selector (Plain: the single selector;
@@ -377,8 +268,8 @@ pub enum K8sDiscoveryMode {
}, },
} }
/// Error returned by [`K8sDiscoveryConfig::mode`] when the selector /// Error returned by [`resolve_mode`] when the selector combination is
/// combination is invalid. /// invalid.
#[derive(Debug, thiserror::Error)] #[derive(Debug, thiserror::Error)]
pub enum ConfigError { 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")]
@@ -391,7 +282,7 @@ pub enum ConfigError {
"discovery.k8s: {selector}_selector `{value}` uses unsupported syntax — \ "discovery.k8s: {selector}_selector `{value}` uses unsupported syntax — \
only equality terms (`key=value` or `key==value`) joined by `,` are accepted. \ only equality terms (`key=value` or `key==value`) joined by `,` are accepted. \
Set-based operators (`in`, `notin`), presence tests, and `!=` silently match \ Set-based operators (`in`, `notin`), presence tests, and `!=` silently match \
zero endpoints at runtime and are rejected at config-load time." zero endpoints at runtime and are rejected at startup."
)] )]
UnsupportedSelectorGrammar { UnsupportedSelectorGrammar {
selector: &'static str, selector: &'static str,
@@ -487,14 +378,16 @@ fn is_equality_selector(selector: &str) -> bool {
true true
} }
impl K8sDiscoveryConfig { /// Validate the selector combination and return the resolved
/// Validate the selector combination and return the resolved mode. /// [`K8sDiscoveryMode`]. Called once at construction by
pub fn mode(&self) -> Result<K8sDiscoveryMode, ConfigError> { /// [`crate::config::cli::Cli::build_discovery`], so an invalid
let plain = self.label_selector.as_deref(); /// combination can never be stored on a [`K8sDiscoveryConfig`].
let prefill = self.prefill_selector.as_deref(); pub fn resolve_mode(
let decode = self.decode_selector.as_deref(); label_selector: Option<&str>,
prefill_selector: Option<&str>,
match (plain, prefill, decode) { decode_selector: Option<&str>,
) -> Result<K8sDiscoveryMode, ConfigError> {
match (label_selector, prefill_selector, decode_selector) {
(Some(label), None, None) => { (Some(label), None, None) => {
// Plain mode pushes `label` to the K8s API as the // Plain mode pushes `label` to the K8s API as the
// server-side `labelSelector` of the EndpointSlice // server-side `labelSelector` of the EndpointSlice
@@ -503,14 +396,12 @@ impl K8sDiscoveryConfig {
// full label-selector grammar — equality, set-based // full label-selector grammar — equality, set-based
// (`in` / `notin`), presence (`key` / `!key`), and // (`in` / `notin`), presence (`key` / `!key`), and
// `!=` — and rejects malformed selectors at // `!=` — and rejects malformed selectors at
// watch-start time. So at config-load we don't // watch-start time. So we don't grammar-check `label`
// grammar-check `label` and let the K8s API be the // here and let the K8s API be the syntax authority. PD
// syntax authority (README.md:25 and the multi-model // mode, in contrast, evaluates selectors client-side via
// e2e in tests/e2e/k8s_integration/test_multi_model.py // `labels_match_selector` which only understands
// depend on this). PD mode, in contrast, evaluates // equality — so PD selectors are still grammar-checked
// selectors client-side via `labels_match_selector` // below.
// which only understands equality — so PD selectors
// are still grammar-checked below.
Ok(K8sDiscoveryMode::Plain { Ok(K8sDiscoveryMode::Plain {
label_selector: label.to_string(), label_selector: label.to_string(),
}) })
@@ -562,29 +453,18 @@ impl K8sDiscoveryConfig {
(Some(_), _, _) => Err(ConfigError::MixedModes), (Some(_), _, _) => Err(ConfigError::MixedModes),
} }
} }
}
#[cfg(test)] #[cfg(test)]
mod k8s_discovery_config_tests { mod k8s_discovery_config_tests {
use super::*; use super::*;
fn cfg(plain: Option<&str>, prefill: Option<&str>, decode: Option<&str>) -> K8sDiscoveryConfig {
K8sDiscoveryConfig {
namespace: "ns".to_string(),
label_selector: plain.map(str::to_string),
prefill_selector: prefill.map(str::to_string),
decode_selector: decode.map(str::to_string),
}
}
#[test] #[test]
fn mode_constructs_pd_disaggregation_from_prefill_and_decode_selectors() { fn mode_constructs_pd_disaggregation_from_prefill_and_decode_selectors() {
// K8s PD now works without per-pod annotations: each worker's // K8s PD now works without per-pod annotations: each worker's
// `/server_info` carries `disaggregation_bootstrap_port`, and the // `/server_info` carries `disaggregation_bootstrap_port`, and the
// worker manager applies it post-discovery. The K8s config layer's // worker manager applies it post-discovery. The K8s config layer's
// job is just to validate the selector combination. // job is just to validate the selector combination.
let m = cfg(None, Some("app=sglang,role=p"), Some("app=sglang,role=d")) let m = resolve_mode(None, Some("app=sglang,role=p"), Some("app=sglang,role=d"))
.mode()
.expect("PD mode is now valid"); .expect("PD mode is now valid");
assert_eq!( assert_eq!(
m, m,
@@ -600,9 +480,8 @@ mod k8s_discovery_config_tests {
// Both PD selectors get the same equality-only grammar check as // Both PD selectors get the same equality-only grammar check as
// the plain label_selector. A set-based prefill selector would // the plain label_selector. A set-based prefill selector would
// silently match zero pods at runtime → fail-fast at load. // silently match zero pods at runtime → fail-fast at load.
let err = cfg(None, Some("app in (sglang, vllm)"), Some("app=sglang")) let err =
.mode() resolve_mode(None, Some("app in (sglang, vllm)"), Some("app=sglang")).unwrap_err();
.unwrap_err();
assert!( assert!(
matches!( matches!(
err, err,
@@ -617,9 +496,8 @@ mod k8s_discovery_config_tests {
#[test] #[test]
fn mode_pd_rejects_set_based_decode_selector() { fn mode_pd_rejects_set_based_decode_selector() {
let err = cfg(None, Some("app=sglang"), Some("app in (sglang, vllm)")) let err =
.mode() resolve_mode(None, Some("app=sglang"), Some("app in (sglang, vllm)")).unwrap_err();
.unwrap_err();
assert!( assert!(
matches!( matches!(
err, err,
@@ -634,7 +512,7 @@ mod k8s_discovery_config_tests {
#[test] #[test]
fn mode_accepts_plain_with_equality_selector() { fn mode_accepts_plain_with_equality_selector() {
let m = cfg(Some("app=sglang"), None, None).mode().unwrap(); let m = resolve_mode(Some("app=sglang"), None, None).unwrap();
assert_eq!( assert_eq!(
m, m,
K8sDiscoveryMode::Plain { K8sDiscoveryMode::Plain {
@@ -646,15 +524,12 @@ mod k8s_discovery_config_tests {
/// Plain mode pushes its selector to the K8s API server-side /// Plain mode pushes its selector to the K8s API server-side
/// (`watcher::Config::default().labels(&selector)` in /// (`watcher::Config::default().labels(&selector)` in
/// `discovery::k8s::spawn`), so the full K8s label-selector grammar /// `discovery::k8s::spawn`), so the full K8s label-selector grammar
/// — including set-based operators — is supported. README.md:25 /// — including set-based operators like `app in (a,b)` — is
/// advertises this, and `tests/e2e/k8s_integration/test_multi_model.py` /// supported and must not be grammar-checked at startup. PD mode
/// relies on it (`label_selector = "app in (sglang,sglang-small)"`). /// (checked client-side) is the opposite; see the PD tests below.
/// Rejecting set-based selectors at config-load broke the documented
/// multi-model k8s path.
#[test] #[test]
fn mode_accepts_set_based_selector_in_plain_mode() { fn mode_accepts_set_based_selector_in_plain_mode() {
let m = cfg(Some("app in (sglang,sglang-small)"), None, None) let m = resolve_mode(Some("app in (sglang,sglang-small)"), None, None)
.mode()
.expect("plain mode must accept set-based selectors"); .expect("plain mode must accept set-based selectors");
assert_eq!( assert_eq!(
m, m,
@@ -675,8 +550,7 @@ mod k8s_discovery_config_tests {
"!deprecated", "!deprecated",
"tier!=canary", "tier!=canary",
] { ] {
let m = cfg(Some(raw), None, None) let m = resolve_mode(Some(raw), None, None)
.mode()
.unwrap_or_else(|e| panic!("plain mode must accept `{raw}`, got {e:?}")); .unwrap_or_else(|e| panic!("plain mode must accept `{raw}`, got {e:?}"));
assert_eq!( assert_eq!(
m, m,
@@ -699,9 +573,8 @@ mod k8s_discovery_config_tests {
/// — both must keep failing. /// — both must keep failing.
#[test] #[test]
fn mode_pd_rejects_notin_prefill_selector() { fn mode_pd_rejects_notin_prefill_selector() {
let err = cfg(None, Some("app notin (vllm, trtllm)"), Some("app=sglang")) let err =
.mode() resolve_mode(None, Some("app notin (vllm, trtllm)"), Some("app=sglang")).unwrap_err();
.unwrap_err();
assert!( assert!(
matches!( matches!(
err, err,
@@ -717,9 +590,7 @@ mod k8s_discovery_config_tests {
#[test] #[test]
fn mode_accepts_comma_separated_equality_terms() { fn mode_accepts_comma_separated_equality_terms() {
// The canonical Plain-mode selector form: `key1=v1,key2=v2`. // The canonical Plain-mode selector form: `key1=v1,key2=v2`.
let m = cfg(Some("app=sglang,zone=us-east"), None, None) let m = resolve_mode(Some("app=sglang,zone=us-east"), None, None).unwrap();
.mode()
.unwrap();
assert_eq!( assert_eq!(
m, m,
K8sDiscoveryMode::Plain { K8sDiscoveryMode::Plain {
@@ -730,30 +601,29 @@ mod k8s_discovery_config_tests {
#[test] #[test]
fn mode_rejects_when_no_selector_is_set() { fn mode_rejects_when_no_selector_is_set() {
let err = cfg(None, None, None).mode().unwrap_err(); let err = resolve_mode(None, None, None).unwrap_err();
assert!(matches!(err, ConfigError::NoSelector), "got {err:?}"); assert!(matches!(err, ConfigError::NoSelector), "got {err:?}");
} }
#[test] #[test]
fn mode_rejects_mixed_plain_and_pd_selectors() { fn mode_rejects_mixed_plain_and_pd_selectors() {
let err = cfg( let err = resolve_mode(
Some("app=sglang"), Some("app=sglang"),
Some("role=prefill"), Some("role=prefill"),
Some("role=decode"), Some("role=decode"),
) )
.mode()
.unwrap_err(); .unwrap_err();
assert!(matches!(err, ConfigError::MixedModes), "got {err:?}"); assert!(matches!(err, ConfigError::MixedModes), "got {err:?}");
} }
#[test] #[test]
fn mode_rejects_partial_pd_selectors() { fn mode_rejects_partial_pd_selectors() {
let err = cfg(None, Some("role=prefill"), None).mode().unwrap_err(); let err = resolve_mode(None, Some("role=prefill"), None).unwrap_err();
assert!( assert!(
matches!(err, ConfigError::PartialPdSelectors), matches!(err, ConfigError::PartialPdSelectors),
"got {err:?}" "got {err:?}"
); );
let err = cfg(None, None, Some("role=decode")).mode().unwrap_err(); let err = resolve_mode(None, None, Some("role=decode")).unwrap_err();
assert!( assert!(
matches!(err, ConfigError::PartialPdSelectors), matches!(err, ConfigError::PartialPdSelectors),
"got {err:?}" "got {err:?}"
@@ -765,7 +635,7 @@ mod k8s_discovery_config_tests {
/// operator opts in by setting plain mode at all). /// operator opts in by setting plain mode at all).
#[test] #[test]
fn mode_accepts_empty_plain_label_selector() { fn mode_accepts_empty_plain_label_selector() {
let m = cfg(Some(""), None, None).mode().unwrap(); let m = resolve_mode(Some(""), None, None).unwrap();
assert_eq!( assert_eq!(
m, m,
K8sDiscoveryMode::Plain { K8sDiscoveryMode::Plain {
@@ -782,7 +652,7 @@ mod k8s_discovery_config_tests {
/// at config load. /// at config load.
#[test] #[test]
fn mode_pd_rejects_empty_prefill_selector() { fn mode_pd_rejects_empty_prefill_selector() {
let err = cfg(None, Some(""), Some("role=decode")).mode().unwrap_err(); let err = resolve_mode(None, Some(""), Some("role=decode")).unwrap_err();
assert!( assert!(
matches!( matches!(
err, err,
@@ -796,9 +666,7 @@ mod k8s_discovery_config_tests {
#[test] #[test]
fn mode_pd_rejects_empty_decode_selector() { fn mode_pd_rejects_empty_decode_selector() {
let err = cfg(None, Some("role=prefill"), Some("")) let err = resolve_mode(None, Some("role=prefill"), Some("")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::EmptyPdSelector { selector: "decode" },), matches!(err, ConfigError::EmptyPdSelector { selector: "decode" },),
"expected EmptyPdSelector(decode), got {err:?}", "expected EmptyPdSelector(decode), got {err:?}",
@@ -810,9 +678,7 @@ mod k8s_discovery_config_tests {
/// failure mode as a literal empty string. /// failure mode as a literal empty string.
#[test] #[test]
fn mode_pd_rejects_whitespace_only_prefill_selector() { fn mode_pd_rejects_whitespace_only_prefill_selector() {
let err = cfg(None, Some(" , "), Some("role=decode")) let err = resolve_mode(None, Some(" , "), Some("role=decode")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!( matches!(
err, err,
@@ -829,9 +695,7 @@ mod k8s_discovery_config_tests {
/// so the decode pool stays empty. /// so the decode pool stays empty.
#[test] #[test]
fn mode_pd_rejects_identical_prefill_and_decode_selectors() { fn mode_pd_rejects_identical_prefill_and_decode_selectors() {
let err = cfg(None, Some("app=sglang"), Some("app=sglang")) let err = resolve_mode(None, Some("app=sglang"), Some("app=sglang")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::IdenticalPdSelectors), matches!(err, ConfigError::IdenticalPdSelectors),
"expected IdenticalPdSelectors, got {err:?}", "expected IdenticalPdSelectors, got {err:?}",
@@ -842,9 +706,7 @@ mod k8s_discovery_config_tests {
/// identical-selector check. /// identical-selector check.
#[test] #[test]
fn mode_pd_rejects_identical_selectors_under_whitespace_normalization() { fn mode_pd_rejects_identical_selectors_under_whitespace_normalization() {
let err = cfg(None, Some("app=sglang"), Some(" app=sglang ")) let err = resolve_mode(None, Some("app=sglang"), Some(" app=sglang ")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::IdenticalPdSelectors), matches!(err, ConfigError::IdenticalPdSelectors),
"expected IdenticalPdSelectors, got {err:?}", "expected IdenticalPdSelectors, got {err:?}",
@@ -861,9 +723,7 @@ mod k8s_discovery_config_tests {
/// string level. /// string level.
#[test] #[test]
fn mode_pd_rejects_identical_selectors_under_eq_alias() { fn mode_pd_rejects_identical_selectors_under_eq_alias() {
let err = cfg(None, Some("app=sglang"), Some("app==sglang")) let err = resolve_mode(None, Some("app=sglang"), Some("app==sglang")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::IdenticalPdSelectors), matches!(err, ConfigError::IdenticalPdSelectors),
"expected IdenticalPdSelectors, got {err:?}", "expected IdenticalPdSelectors, got {err:?}",
@@ -877,9 +737,7 @@ mod k8s_discovery_config_tests {
/// form must agree. /// form must agree.
#[test] #[test]
fn mode_pd_rejects_identical_selectors_under_inner_whitespace() { fn mode_pd_rejects_identical_selectors_under_inner_whitespace() {
let err = cfg(None, Some("app=sglang"), Some("app = sglang")) let err = resolve_mode(None, Some("app=sglang"), Some("app = sglang")).unwrap_err();
.mode()
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::IdenticalPdSelectors), matches!(err, ConfigError::IdenticalPdSelectors),
"expected IdenticalPdSelectors, got {err:?}", "expected IdenticalPdSelectors, got {err:?}",
@@ -893,9 +751,8 @@ mod k8s_discovery_config_tests {
/// reintroduce the silent-failure bug.) /// reintroduce the silent-failure bug.)
#[test] #[test]
fn mode_pd_rejects_identical_selectors_under_term_order_permutation() { fn mode_pd_rejects_identical_selectors_under_term_order_permutation() {
let err = cfg(None, Some("role=p,app=sglang"), Some("app=sglang,role=p")) let err =
.mode() resolve_mode(None, Some("role=p,app=sglang"), Some("app=sglang,role=p")).unwrap_err();
.unwrap_err();
assert!( assert!(
matches!(err, ConfigError::IdenticalPdSelectors), matches!(err, ConfigError::IdenticalPdSelectors),
"expected IdenticalPdSelectors, got {err:?}", "expected IdenticalPdSelectors, got {err:?}",
@@ -907,12 +764,11 @@ mod k8s_discovery_config_tests {
/// aggressive that it false-positives on legitimate PD configs. /// aggressive that it false-positives on legitimate PD configs.
#[test] #[test]
fn mode_pd_accepts_truly_distinct_selectors() { fn mode_pd_accepts_truly_distinct_selectors() {
let m = cfg( let m = resolve_mode(
None, None,
Some("app=sglang,role=prefill"), Some("app=sglang,role=prefill"),
Some("app=sglang,role=decode"), Some("app=sglang,role=decode"),
) )
.mode()
.expect("distinct selectors must validate"); .expect("distinct selectors must validate");
assert!(matches!(m, K8sDiscoveryMode::PdDisaggregation { .. })); assert!(matches!(m, K8sDiscoveryMode::PdDisaggregation { .. }));
} }
+35 -3
View File
@@ -328,16 +328,18 @@ pub async fn spawn(
cfg: K8sDiscoveryConfig, cfg: K8sDiscoveryConfig,
tx: mpsc::Sender<DiscoveryEvent>, tx: mpsc::Sender<DiscoveryEvent>,
) -> Result<tokio::task::JoinHandle<()>> { ) -> Result<tokio::task::JoinHandle<()>> {
let mode = cfg.mode().context("validate k8s discovery selectors")?; // The mode was resolved + validated at construction (`resolve_mode` in
// `Cli::build_discovery`); just destructure it here.
let K8sDiscoveryConfig { namespace, mode } = cfg;
let client = Client::try_default() let client = Client::try_default()
.await .await
.context("kube client default config")?; .context("kube client default config")?;
let api: Api<EndpointSlice> = if cfg.namespace.is_empty() { let api: Api<EndpointSlice> = if namespace.is_empty() {
Api::all(client) Api::all(client)
} else { } else {
Api::namespaced(client, &cfg.namespace) Api::namespaced(client, &namespace)
}; };
// Plain mode pushes the single selector to the server side so the LIST // Plain mode pushes the single selector to the server side so the LIST
@@ -351,6 +353,36 @@ pub async fn spawn(
}; };
let watcher_cfg = watcher::Config::default().labels(&server_side_selector); let watcher_cfg = watcher::Config::default().labels(&server_side_selector);
// Log the resolved namespace + selector(s) at startup. We can't
// verify the namespace exists (the router's RBAC covers
// endpointslices/services/pods, not namespaces, and a correct
// namespace legitimately has zero matching workers until they come
// up), so a typo'd `--service-discovery-namespace` silently watches
// an empty namespace. Surfacing the watch target here lets an
// operator spot the typo in the first log lines instead of only
// discovering it via later `no workers available` request failures.
let namespace_display: &str = if namespace.is_empty() {
"<all namespaces>"
} else {
&namespace
};
match &mode {
K8sDiscoveryMode::Plain { label_selector } => tracing::info!(
namespace = %namespace_display,
label_selector = %label_selector,
"k8s discovery starting (plain mode); a wrong namespace or selector matches zero EndpointSlices"
),
K8sDiscoveryMode::PdDisaggregation {
prefill_selector,
decode_selector,
} => tracing::info!(
namespace = %namespace_display,
prefill_selector = %prefill_selector,
decode_selector = %decode_selector,
"k8s discovery starting (PD mode); a wrong namespace or selector matches zero EndpointSlices"
),
}
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
let stream = watcher(api, watcher_cfg); let stream = watcher(api, watcher_cfg);
tokio::pin!(stream); tokio::pin!(stream);
+1 -1
View File
@@ -25,7 +25,7 @@ pub async fn spawn_discovery(
cfg: &Config, cfg: &Config,
) -> Result<(mpsc::Receiver<DiscoveryEvent>, tokio::task::JoinHandle<()>)> { ) -> Result<(mpsc::Receiver<DiscoveryEvent>, tokio::task::JoinHandle<()>)> {
let (tx, rx) = mpsc::channel(DISCOVERY_CHANNEL_CAP); let (tx, rx) = mpsc::channel(DISCOVERY_CHANNEL_CAP);
let handle = match &cfg.discovery.backend { let handle = match &cfg.discovery {
DiscoveryBackend::StaticUrls(s) => static_urls::spawn(s.clone(), tx).await?, DiscoveryBackend::StaticUrls(s) => static_urls::spawn(s.clone(), tx).await?,
DiscoveryBackend::K8s(k) => k8s::spawn(k.clone(), tx).await?, DiscoveryBackend::K8s(k) => k8s::spawn(k.clone(), tx).await?,
}; };
+9 -16
View File
@@ -3,18 +3,10 @@
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use clap::Parser; use clap::Parser;
use sgl_router::config::LogFormat; use sgl_router::config::{Cli, LogFormat};
use std::path::PathBuf;
use std::sync::Arc; use std::sync::Arc;
use tokio::signal::unix::{signal, Signal, SignalKind}; use tokio::signal::unix::{signal, Signal, SignalKind};
#[derive(Parser, Debug)]
#[command(name = "sgl-router", version)]
struct Cli {
#[arg(long, env = "SGL_ROUTER_CONFIG")]
config: PathBuf,
}
/// Install the global tracing subscriber. /// Install the global tracing subscriber.
/// ///
/// Idempotent: a second call returns `Ok` without panicking. When /// Idempotent: a second call returns `Ok` without panicking. When
@@ -54,13 +46,13 @@ fn init_tracing(default_level: &str, format: LogFormat) -> Result<()> {
Ok(()) Ok(())
} }
/// Install a minimal text-format subscriber BEFORE config parsing so a /// Install a minimal text-format subscriber BEFORE config resolution so a
/// config-load error has somewhere to surface. The real subscriber /// config-resolution error has somewhere to surface. The real subscriber
/// (driven by `Config.observability`) is installed after; the second /// (driven by `Config.observability`) is installed after; the second
/// `try_init` is a no-op because a subscriber is already present. /// `try_init` is a no-op because a subscriber is already present.
/// The bootstrap subscriber respects `RUST_LOG` so an operator can /// The bootstrap subscriber respects `RUST_LOG` so an operator can
/// debug startup with `RUST_LOG=debug` even when the config file is /// debug startup with `RUST_LOG=debug` even when configuration resolution
/// missing or malformed. /// fails.
fn install_bootstrap_subscriber() { fn install_bootstrap_subscriber() {
let filter = tracing_subscriber::EnvFilter::try_from_default_env() let filter = tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")); .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info"));
@@ -83,12 +75,13 @@ fn install_signal_handlers() -> Result<(Signal, Signal)> {
#[tokio::main] #[tokio::main]
async fn main() -> Result<()> { async fn main() -> Result<()> {
let cli = Cli::parse(); let cli = Cli::parse();
// Bootstrap subscriber so a Config::from_path error has structured // Bootstrap subscriber so a config-resolution error has structured
// output. The configured-format subscriber installs after this and // output. The configured-format subscriber installs after this and
// becomes a no-op via try_init's idempotency. // becomes a no-op via try_init's idempotency.
install_bootstrap_subscriber(); install_bootstrap_subscriber();
let cfg = sgl_router::config::Config::from_path(&cli.config) let cfg = cli
.with_context(|| format!("load config from {}", cli.config.display()))?; .into_config()
.context("resolve configuration from CLI flags")?;
init_tracing(&cfg.observability.log_level, cfg.observability.log_format)?; init_tracing(&cfg.observability.log_level, cfg.observability.log_format)?;
@@ -36,11 +36,14 @@
use crate::config::CacheAwareConfig; use crate::config::CacheAwareConfig;
use crate::discovery::ModelId; use crate::discovery::ModelId;
use crate::policies::kv_events::{compute_block_hashes, BlockSizeOracle, HashTree}; use crate::policies::kv_events::{
compute_block_hashes, compute_block_hashes_bigram, BlockSizeOracle, HashTree,
};
use crate::policies::{Policy, SelectionContext}; use crate::policies::{Policy, SelectionContext};
use crate::server::metrics::MetricsRegistry;
use crate::tokenizer::{adapter, TokenizerRegistry}; use crate::tokenizer::{adapter, TokenizerRegistry};
use crate::workers::Worker; use crate::workers::Worker;
use std::sync::Arc; use std::sync::{Arc, OnceLock};
/// Selection policy that scores candidates by tree-overlap with the /// Selection policy that scores candidates by tree-overlap with the
/// request's prefix and falls back to load-based picking when the tree /// request's prefix and falls back to load-based picking when the tree
@@ -59,6 +62,14 @@ pub struct CacheAwareZmqPolicy {
/// degrades to min-load — the router cannot hash a prompt without /// degrades to min-load — the router cannot hash a prompt without
/// a block size that matches what the worker publishes. /// a block size that matches what the worker publishes.
block_size_oracle: Arc<BlockSizeOracle>, block_size_oracle: Arc<BlockSizeOracle>,
/// Optional metrics sink. Set via [`Self::with_metrics`] by the policy
/// factory for the production policy; `None` in unit tests and
/// non-cache-aware call sites. When set, each cache-aware selection
/// records the prefix-overlap block count into
/// `sgl_router_overlap_blocks`. Set once via [`Self::with_metrics`]
/// (tests) or the `Policy::attach_metrics` hook (production, called by
/// `PolicyRegistry::attach_metrics` after the registry is built).
metrics: OnceLock<Arc<MetricsRegistry>>,
} }
impl std::fmt::Debug for CacheAwareZmqPolicy { impl std::fmt::Debug for CacheAwareZmqPolicy {
@@ -82,9 +93,19 @@ impl CacheAwareZmqPolicy {
tree, tree,
tokenizers, tokenizers,
block_size_oracle, block_size_oracle,
metrics: OnceLock::new(),
} }
} }
/// Attach a metrics sink so each cache-aware selection records the
/// prefix-overlap block count into `sgl_router_overlap_blocks`. Builder
/// form used by tests; production wiring goes through the
/// `Policy::attach_metrics` hook.
pub fn with_metrics(self, metrics: Arc<MetricsRegistry>) -> Self {
let _ = self.metrics.set(metrics);
self
}
/// Lowest-load worker — ties broken by stable iteration order (which /// Lowest-load worker — ties broken by stable iteration order (which
/// is the order the registry returned, i.e. dashmap-undefined). For /// is the order the registry returned, i.e. dashmap-undefined). For
/// production traffic the ties are rare; tests pin the load skew. /// production traffic the ties are rare; tests pin the load skew.
@@ -218,9 +239,22 @@ impl Policy for CacheAwareZmqPolicy {
// has registered yet (oracle empty), cache-aware routing has no // has registered yet (oracle empty), cache-aware routing has no
// ground truth to score against; fall back to min-load. // ground truth to score against; fall back to min-load.
let Some(block_size) = self.block_size_oracle.get() else { let Some(block_size) = self.block_size_oracle.get() else {
tracing::debug!(
model = %ctx.model(),
"cache-aware-zmq: block size unknown (no worker page_size yet), falling back to min-load",
);
return Self::pick_min_load(workers); return Self::pick_min_load(workers);
}; };
let block_hashes = compute_block_hashes(&tokens, block_size as usize); // EAGLE-family workers hash KV blocks over token bigrams; the query
// hashes must match the worker's stored hashes or the tree lookup
// always misses (overlap stays 0). The oracle carries the worker-
// reported flag.
let is_bigram = self.block_size_oracle.is_bigram();
let block_hashes = if is_bigram {
compute_block_hashes_bigram(&tokens, block_size as usize)
} else {
compute_block_hashes(&tokens, block_size as usize)
};
if block_hashes.is_empty() { if block_hashes.is_empty() {
return Self::pick_min_load(workers); return Self::pick_min_load(workers);
} }
@@ -228,13 +262,28 @@ impl Policy for CacheAwareZmqPolicy {
let match_rate = matched.matched_blocks as f32 / block_hashes.len() as f32; let match_rate = matched.matched_blocks as f32 / block_hashes.len() as f32;
tracing::debug!( tracing::debug!(
model = %ctx.model(), model = %ctx.model(),
hashing = if is_bigram { "bigram" } else { "unigram" },
n_blocks = block_hashes.len(), n_blocks = block_hashes.len(),
matched_blocks = matched.matched_blocks, matched_blocks = matched.matched_blocks,
match_rate, match_rate,
cache_threshold = self.config.cache_threshold, cache_threshold = self.config.cache_threshold,
"cache-aware-zmq match_prefix", "cache-aware-zmq match_prefix",
); );
// Record the matched overlap into `sgl_router_overlap_blocks` before
// the threshold branch, so the histogram captures the full
// distribution — including low-overlap selections that fall back to
// min-load. This is the quantitative signal that cache-aware routing
// is matching prefixes at all.
if let Some(m) = self.metrics.get() {
m.observe_overlap_blocks(ctx.model().0.as_str(), matched.matched_blocks as u64);
}
if match_rate <= self.config.cache_threshold || matched.workers.is_empty() { if match_rate <= self.config.cache_threshold || matched.workers.is_empty() {
tracing::debug!(
model = %ctx.model(),
match_rate,
cache_threshold = self.config.cache_threshold,
"cache-aware-zmq: overlap below threshold, falling back to min-load",
);
return Self::pick_min_load(workers); return Self::pick_min_load(workers);
} }
// Among workers in the matched set, pick the lowest-load one. // Among workers in the matched set, pick the lowest-load one.
@@ -245,7 +294,20 @@ impl Policy for CacheAwareZmqPolicy {
.filter(|w| matched_urls.contains(w.url.as_str())) .filter(|w| matched_urls.contains(w.url.as_str()))
.min_by_key(|w| w.active_load()) .min_by_key(|w| w.active_load())
.map(Arc::clone); .map(Arc::clone);
best_matched.or_else(|| Self::pick_min_load(workers)) let chosen = best_matched.or_else(|| Self::pick_min_load(workers));
if let Some(w) = &chosen {
tracing::debug!(
model = %ctx.model(),
worker = %w.url,
matched_blocks = matched.matched_blocks,
"cache-aware-zmq: selected worker by cache overlap",
);
}
chosen
}
fn attach_metrics(&self, metrics: Arc<MetricsRegistry>) {
let _ = self.metrics.set(metrics);
} }
} }
@@ -292,20 +354,18 @@ mod tests {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![crate::config::ModelConfig { model: crate::config::ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: crate::config::PolicyKind::RoundRobin, policy: crate::config::PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: crate::config::DiscoveryConfig { discovery: crate::config::DiscoveryBackend::StaticUrls(
backend: crate::config::DiscoveryBackend::StaticUrls(
crate::config::StaticUrlsDiscoveryConfig { crate::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}, },
), ),
},
proxy: crate::config::ProxyConfig::default(), proxy: crate::config::ProxyConfig::default(),
active_load: crate::config::ActiveLoadConfig::default(), active_load: crate::config::ActiveLoadConfig::default(),
}; };
@@ -392,6 +452,266 @@ mod tests {
assert_eq!(chosen.url, "http://w0:30000"); assert_eq!(chosen.url, "http://w0:30000");
} }
/// The cache-aware path records the matched prefix-overlap block count
/// into `sgl_router_overlap_blocks`. Regression: the metric was defined
/// but never observed in production, so the histogram stayed empty and
/// gave no signal that cache-aware routing was matching anything.
#[test]
fn records_overlap_blocks_metric() {
let tree = Arc::new(HashTree::new());
let registry = tokenizer_registry_with_tiny();
let text = "hello world hello world hello world";
let tok = registry.get("tiny").unwrap();
let ids = adapter::encode(&tok, text).unwrap();
let block_size = 4u32;
let hashes = compute_block_hashes(&ids, block_size as usize);
assert!(!hashes.is_empty());
tree.insert(&KvWorkerId::new("http://w0:30000".into(), 0), None, &hashes);
let metrics = MetricsRegistry::new();
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 0.0,
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
},
tree,
registry,
oracle_for_tests(4),
)
.with_metrics(Arc::clone(&metrics));
let workers = vec![
worker("http://w0:30000", "tiny"),
worker("http://w1:30000", "tiny"),
];
let model = ModelId("tiny".into());
let body = serde_json::to_vec(&serde_json::json!({"prompt": text})).unwrap();
let ctx = SelectionContext::new(&model, Some(&body));
let _ = policy.select(&workers, &ctx).expect("must pick");
let rendered = metrics.render();
assert!(
rendered.contains("sgl_router_overlap_blocks_count{model_id=\"tiny\"}"),
"overlap_blocks histogram must be observed on a cache-aware selection; got:\n{rendered}"
);
}
/// Production wiring path: the policy is stored as `Arc<dyn Policy>` in a
/// `PolicyRegistry`, then `PolicyRegistry::attach_metrics` injects the
/// registry — exactly what `AppContext::with_active_load` does at startup.
/// Exercises trait dispatch (the default no-op vs the `CacheAwareZmqPolicy`
/// override) and the registry fan-out, neither of which the `with_metrics`
/// builder test covers.
#[test]
fn attach_metrics_via_registry_records_overlap() {
let tree = Arc::new(HashTree::new());
let toks = tokenizer_registry_with_tiny();
let text = "hello world hello world hello world";
let tok = toks.get("tiny").unwrap();
let ids = adapter::encode(&tok, text).unwrap();
let hashes = compute_block_hashes(&ids, 4);
assert!(!hashes.is_empty());
tree.insert(&KvWorkerId::new("http://w0:30000".into(), 0), None, &hashes);
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 0.0,
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
},
tree,
toks,
oracle_for_tests(4),
);
let model = ModelId("tiny".into());
let registry = crate::policies::PolicyRegistry::default();
registry.insert(model.clone(), Arc::new(policy));
// The production injection point — not the `with_metrics` builder.
let metrics = MetricsRegistry::new();
registry.attach_metrics(Arc::clone(&metrics));
let chosen_policy = registry.get(&model).unwrap();
let workers = vec![
worker("http://w0:30000", "tiny"),
worker("http://w1:30000", "tiny"),
];
let body = serde_json::to_vec(&serde_json::json!({"prompt": text})).unwrap();
let ctx = SelectionContext::new(&model, Some(&body));
let _ = chosen_policy.select(&workers, &ctx).expect("must pick");
let rendered = metrics.render();
assert!(
rendered.contains("sgl_router_overlap_blocks_count{model_id=\"tiny\"}"),
"PolicyRegistry::attach_metrics must wire overlap recording through the trait; got:\n{rendered}"
);
}
/// The overlap observation is recorded *before* the cache-threshold branch,
/// so low-overlap selections that fall back to min-load are still counted.
/// `cache_threshold: 1.0` forces the fallback (match_rate is always <= 1.0)
/// even on a full prefix match; assert the histogram is still observed AND
/// the pick came from min-load (w1), not the cache-overlap worker (w0).
#[test]
fn overlap_recorded_even_when_selection_falls_back() {
let tree = Arc::new(HashTree::new());
let toks = tokenizer_registry_with_tiny();
let text = "hello world hello world hello world";
let tok = toks.get("tiny").unwrap();
let ids = adapter::encode(&tok, text).unwrap();
let hashes = compute_block_hashes(&ids, 4);
assert!(!hashes.is_empty());
tree.insert(&KvWorkerId::new("http://w0:30000".into(), 0), None, &hashes);
let metrics = MetricsRegistry::new();
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 1.0, // match_rate <= 1.0 always -> always fall back
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
},
tree,
toks,
oracle_for_tests(4),
)
.with_metrics(Arc::clone(&metrics));
// Bump w0's load so min-load picks w1 — distinguishing a min-load
// fallback from the cache-overlap pick (which would be w0). Two guards
// mirror `empty_tree_falls_back_to_min_load` (below the imbalance
// threshold, so the cache-aware path is still reached).
let w0 = worker("http://w0:30000", "tiny");
let w1 = worker("http://w1:30000", "tiny");
let _g = w0.load_guard();
let _g2 = w0.load_guard();
let workers = vec![Arc::clone(&w0), Arc::clone(&w1)];
let model = ModelId("tiny".into());
let body = serde_json::to_vec(&serde_json::json!({"prompt": text})).unwrap();
let ctx = SelectionContext::new(&model, Some(&body));
let chosen = policy.select(&workers, &ctx).expect("must pick");
assert_eq!(
chosen.url, "http://w1:30000",
"cache_threshold 1.0 must force a min-load fallback (w1), not the overlap worker (w0)"
);
let rendered = metrics.render();
assert!(
rendered.contains("sgl_router_overlap_blocks_count{model_id=\"tiny\"}"),
"overlap must be recorded even on the below-threshold fallback; got:\n{rendered}"
);
}
/// End-to-end bigram wiring (the fix that takes `overlap_blocks_sum` from
/// 0 to non-zero for EAGLE models): an EAGLE worker publishes its blocks
/// under BIGRAM hashes. Only a router whose oracle reports `is_bigram` —
/// and thus hashes its query with the bigram hasher — matches them, so
/// overlap is non-zero and it picks the cached worker. A unigram-hashing
/// router against the SAME tree matches nothing (overlap recorded as 0).
#[test]
fn bigram_routing_matches_only_with_bigram_hashing() {
fn overlap_sum(rendered: &str) -> f64 {
rendered
.lines()
.find(|l| l.starts_with("sgl_router_overlap_blocks_sum{model_id=\"tiny\"}"))
.and_then(|l| l.split_whitespace().last())
.and_then(|v| v.parse::<f64>().ok())
.unwrap_or(-1.0)
}
let registry = tokenizer_registry_with_tiny();
let text = "hello world hello world hello world";
let tok = registry.get("tiny").unwrap();
let ids = adapter::encode(&tok, text).unwrap();
let block_size = 4u32;
// The EAGLE worker publishes BIGRAM block hashes.
let bigram_hashes = compute_block_hashes_bigram(&ids, block_size as usize);
assert!(!bigram_hashes.is_empty());
assert_ne!(
bigram_hashes,
compute_block_hashes(&ids, block_size as usize),
"bigram and unigram hashes must differ for this prefix"
);
let model = ModelId("tiny".into());
let body = serde_json::to_vec(&serde_json::json!({ "prompt": text })).unwrap();
// Bigram-aware router (oracle.is_bigram == true): query hashes match
// the bigram tree -> overlap > 0 and it picks the matched worker w0.
{
let tree = Arc::new(HashTree::new());
tree.insert(
&KvWorkerId::new("http://w0:30000".into(), 0),
None,
&bigram_hashes,
);
let oracle = BlockSizeOracle::new();
oracle.try_set(block_size).unwrap();
oracle.set_bigram(true);
let metrics = MetricsRegistry::new();
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 0.0,
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
},
tree,
Arc::clone(&registry),
oracle,
)
.with_metrics(Arc::clone(&metrics));
let workers = vec![
worker("http://w0:30000", "tiny"),
worker("http://w1:30000", "tiny"),
];
let ctx = SelectionContext::new(&model, Some(&body));
let chosen = policy.select(&workers, &ctx).expect("must pick");
assert_eq!(
chosen.url, "http://w0:30000",
"bigram-aware router must match w0's bigram-hashed prefix"
);
assert!(
overlap_sum(&metrics.render()) > 0.0,
"overlap_blocks_sum must be > 0 once the router hashes with bigram"
);
}
// Unigram router (default is_bigram == false) vs the SAME bigram tree:
// query hashes never match -> overlap recorded as 0.
{
let tree = Arc::new(HashTree::new());
tree.insert(
&KvWorkerId::new("http://w0:30000".into(), 0),
None,
&bigram_hashes,
);
let oracle = BlockSizeOracle::new();
oracle.try_set(block_size).unwrap();
let metrics = MetricsRegistry::new();
let policy = CacheAwareZmqPolicy::new(
CacheAwareConfig {
cache_threshold: 0.0,
balance_abs_threshold: 32,
balance_rel_threshold: 1.1,
},
tree,
Arc::clone(&registry),
oracle,
)
.with_metrics(Arc::clone(&metrics));
let workers = vec![
worker("http://w0:30000", "tiny"),
worker("http://w1:30000", "tiny"),
];
let ctx = SelectionContext::new(&model, Some(&body));
let _ = policy.select(&workers, &ctx).expect("must pick");
assert_eq!(
overlap_sum(&metrics.render()),
0.0,
"unigram hashing matches nothing in a bigram tree -> overlap_sum == 0"
);
}
}
/// Two workers both hold the prefix; the lower-load one wins. /// Two workers both hold the prefix; the lower-load one wins.
#[test] #[test]
fn tie_break_by_lowest_active_load() { fn tie_break_by_lowest_active_load() {
+12 -22
View File
@@ -76,7 +76,7 @@ pub fn build_registry(
block_size_oracle: Arc<BlockSizeOracle>, block_size_oracle: Arc<BlockSizeOracle>,
) -> Result<PolicyRegistry> { ) -> Result<PolicyRegistry> {
let reg = PolicyRegistry::default(); let reg = PolicyRegistry::default();
for m in &cfg.models { let m = &cfg.model;
reg.insert( reg.insert(
ModelId(m.id.clone()), ModelId(m.id.clone()),
build_policy( build_policy(
@@ -86,7 +86,6 @@ pub fn build_registry(
Arc::clone(&block_size_oracle), Arc::clone(&block_size_oracle),
), ),
); );
}
Ok(reg) Ok(reg)
} }
@@ -111,34 +110,29 @@ pub fn build_registry_with_defaults(cfg: &Config) -> Result<PolicyRegistry> {
mod tests { mod tests {
use super::*; use super::*;
use crate::config::{ use crate::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ProxyConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ProxyConfig, ServerConfig,
ServerConfig, StaticUrlsDiscoveryConfig, StaticUrlsDiscoveryConfig,
}; };
use crate::config::PolicyKind; use crate::config::PolicyKind;
fn cfg_with_models(policies: &[(&str, PolicyKind)]) -> Config { fn cfg_with_model(id: &str, policy: PolicyKind) -> Config {
Config { Config {
server: ServerConfig { server: ServerConfig {
host: "0".into(), host: "0".into(),
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: policies model: ModelConfig {
.iter() id: id.into(),
.map(|(id, p)| ModelConfig {
id: (*id).into(),
tokenizer_path: "/tmp/x".into(), tokenizer_path: "/tmp/x".into(),
policy: *p, policy,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}) },
.collect(), discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
discovery: DiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }
@@ -154,22 +148,18 @@ mod tests {
} }
#[test] #[test]
fn registry_assigns_per_model() { fn registry_assigns_configured_model() {
let cfg = cfg_with_models(&[ let cfg = cfg_with_model("qwen", PolicyKind::RoundRobin);
("qwen", PolicyKind::RoundRobin),
("deepseek", PolicyKind::Random),
]);
let tree = Arc::new(HashTree::new()); let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default()); let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(&cfg, tree, tokenizers, BlockSizeOracle::new()).unwrap(); let reg = build_registry(&cfg, tree, tokenizers, BlockSizeOracle::new()).unwrap();
assert!(reg.get(&ModelId("qwen".into())).is_some()); assert!(reg.get(&ModelId("qwen".into())).is_some());
assert!(reg.get(&ModelId("deepseek".into())).is_some());
assert!(reg.get(&ModelId("missing".into())).is_none()); assert!(reg.get(&ModelId("missing".into())).is_none());
} }
#[test] #[test]
fn cache_aware_zmq_builds_via_factory() { fn cache_aware_zmq_builds_via_factory() {
let cfg = cfg_with_models(&[("modelA", PolicyKind::CacheAwareZmq)]); let cfg = cfg_with_model("modelA", PolicyKind::CacheAwareZmq);
let tree = Arc::new(HashTree::new()); let tree = Arc::new(HashTree::new());
let tokenizers = Arc::new(TokenizerRegistry::default()); let tokenizers = Arc::new(TokenizerRegistry::default());
let reg = build_registry(&cfg, tree, tokenizers, BlockSizeOracle::new()).unwrap(); let reg = build_registry(&cfg, tree, tokenizers, BlockSizeOracle::new()).unwrap();
@@ -31,17 +31,32 @@
//! through `KvEventIndex::add_worker`; that refactor can land later //! through `KvEventIndex::add_worker`; that refactor can land later
//! without changing the oracle's public surface. //! without changing the oracle's public surface.
use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::atomic::{AtomicU32, AtomicU8, Ordering};
use std::sync::Arc; use std::sync::Arc;
/// Tri-state for the bigram flag: distinguishes "not yet reported" from an
/// established `false`, so [`BlockSizeOracle::set_bigram`] can be first-wins
/// (matching `try_set`) rather than last-writer-wins.
const BIGRAM_UNKNOWN: u8 = 0;
const BIGRAM_UNIGRAM: u8 = 1;
const BIGRAM_BIGRAM: u8 = 2;
/// First-wins, idempotent block-size publisher. /// First-wins, idempotent block-size publisher.
/// ///
/// Internally an `AtomicU32` where 0 means "not yet known". Use /// Internally an `AtomicU32` where 0 means "not yet known". Use
/// [`Self::try_set`] to publish a worker-reported value and /// [`Self::try_set`] to publish a worker-reported value and
/// [`Self::get`] to read at routing time. /// [`Self::get`] to read at routing time.
///
/// Also carries a `bigram` flag — EAGLE-family workers hash KV blocks over
/// token bigrams, so the policy must pick the bigram hasher. Like `value` it
/// is a per-cluster property (all workers run the same model) and is
/// established first-wins with a loud warning on disagreement, mirroring
/// `try_set` — a heterogeneous EAGLE/non-EAGLE cluster would otherwise let the
/// last registrant silently flip the global hashing mode.
#[derive(Debug, Default)] #[derive(Debug, Default)]
pub struct BlockSizeOracle { pub struct BlockSizeOracle {
value: AtomicU32, value: AtomicU32,
bigram: AtomicU8,
} }
/// Returned by [`BlockSizeOracle::try_set`] when the candidate disagrees /// Returned by [`BlockSizeOracle::try_set`] when the candidate disagrees
@@ -70,6 +85,46 @@ impl BlockSizeOracle {
} }
} }
/// Publish whether the cluster's workers use bigram (EAGLE-family) KV-block
/// hashing. Called from `KvEventIndex::add_worker` alongside `try_set`.
/// First-wins: the first worker establishes the mode; a later worker that
/// disagrees is logged (not silently honored), since the query-hashing mode
/// is process-wide and one mismatched worker would zero out cache-aware
/// routing for the cluster.
pub fn set_bigram(&self, is_bigram: bool) {
let candidate = if is_bigram {
BIGRAM_BIGRAM
} else {
BIGRAM_UNIGRAM
};
match self.bigram.compare_exchange(
BIGRAM_UNKNOWN,
candidate,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => {}
Err(existing) if existing == candidate => {}
Err(existing) => {
tracing::warn!(
established_bigram = existing == BIGRAM_BIGRAM,
worker_bigram = is_bigram,
"kv-events: worker hashing mode (bigram/EAGLE) disagrees with the \
established cluster value; keeping the first. A heterogeneous \
EAGLE/non-EAGLE cluster will silently never match cache for the \
minority workers — check that all workers run the same model.",
);
}
}
}
/// Whether query hashing should use the bigram variant
/// ([`super::hash::compute_block_hashes_bigram`]). Defaults to `false`
/// until a worker reports an EAGLE-family `speculative_algorithm`.
pub fn is_bigram(&self) -> bool {
self.bigram.load(Ordering::Relaxed) == BIGRAM_BIGRAM
}
/// Publish a candidate block size. Returns the established value on /// Publish a candidate block size. Returns the established value on
/// success (idempotent: same candidate as already set is `Ok`); /// success (idempotent: same candidate as already set is `Ok`);
/// returns `Err(BlockSizeMismatch)` when the candidate disagrees. /// returns `Err(BlockSizeMismatch)` when the candidate disagrees.
@@ -123,6 +178,36 @@ mod tests {
assert_eq!(oracle.get(), Some(64)); assert_eq!(oracle.get(), Some(64));
} }
#[test]
fn bigram_flag_defaults_false_first_wins_and_is_idempotent() {
let oracle = BlockSizeOracle::new();
assert!(
!oracle.is_bigram(),
"unknown (no worker reported yet) reads as non-bigram"
);
oracle.set_bigram(true);
assert!(oracle.is_bigram(), "first worker establishes the mode");
oracle.set_bigram(true); // idempotent agreement
assert!(oracle.is_bigram());
// Independent of block_size establishment.
assert_eq!(oracle.get(), None);
// First-wins: a conflicting later worker is logged, not honored.
oracle.set_bigram(false);
assert!(
oracle.is_bigram(),
"a disagreeing worker must not flip the established mode"
);
}
#[test]
fn bigram_flag_establishes_false_first_wins() {
let oracle = BlockSizeOracle::new();
oracle.set_bigram(false);
assert!(!oracle.is_bigram(), "established as unigram");
oracle.set_bigram(true); // conflicting; first (unigram) wins
assert!(!oracle.is_bigram());
}
#[test] #[test]
fn mismatching_set_fails_without_changing_state() { fn mismatching_set_fails_without_changing_state() {
let oracle = BlockSizeOracle::new(); let oracle = BlockSizeOracle::new();
@@ -48,6 +48,14 @@ pub struct EventConfig {
/// many SUB connections (one per rank), skipping any rank whose /// many SUB connections (one per rank), skipping any rank whose
/// `port_base + dp_rank` overflows `u16`. /// `port_base + dp_rank` overflows `u16`.
pub dp_size: u32, pub dp_size: u32,
/// Whether the worker uses EAGLE-family speculative decoding (EAGLE /
/// EAGLE3 / FROZEN_KV_MTP), reported via `/server_info`'s top-level
/// `speculative_algorithm`. When true the worker hashes KV blocks over
/// overlapping token *bigrams* (`is_bigram = is_eagle`), so the router must
/// use [`super::hash::compute_block_hashes_bigram`] for its query hashes to
/// match the worker's stored hashes — otherwise cache-aware routing
/// silently never matches and degrades to min-load.
pub is_bigram: bool,
} }
/// Default timeout for the `/server_info` introspection request. The /// Default timeout for the `/server_info` introspection request. The
@@ -88,6 +96,10 @@ pub async fn fetch_event_config(
let body = fetch_with_retry(&server_info_url, worker_url, client).await?; let body = fetch_with_retry(&server_info_url, worker_url, client).await?;
// EAGLE-family speculative decoding ⇒ the worker hashes KV blocks over
// token bigrams; the router must mirror that on the selection side.
let is_bigram = classify_bigram(body.speculative_algorithm.as_deref(), worker_url);
let block = match body.kv_events { let block = match body.kv_events {
Some(b) => b, Some(b) => b,
None => { None => {
@@ -117,6 +129,7 @@ pub async fn fetch_event_config(
topic: block.topic, topic: block.topic,
block_size: block.block_size, block_size: block.block_size,
dp_size: block.dp_size, dp_size: block.dp_size,
is_bigram,
})) }))
} }
@@ -185,6 +198,43 @@ async fn fetch_with_retry(
struct ServerInfoResponse { struct ServerInfoResponse {
#[serde(default)] #[serde(default)]
kv_events: Option<KvEventsBlock>, kv_events: Option<KvEventsBlock>,
/// Top-level `/server_info` field. EAGLE-family values
/// (EAGLE / EAGLE3 / FROZEN_KV_MTP) mean the worker hashes KV blocks over
/// token bigrams — see [`EventConfig::is_bigram`].
#[serde(default)]
speculative_algorithm: Option<String>,
}
/// Whether a worker's `/server_info` `speculative_algorithm` means it hashes KV
/// blocks over token bigrams. Recognizes the engine's `is_eagle()` set
/// (`EAGLE`, `EAGLE3`, `FROZEN_KV_MTP`, case-insensitive).
///
/// An *unrecognized* value that looks EAGLE-family (contains `EAGLE` or `MTP`)
/// is logged loudly and treated as non-bigram — it most likely means a new
/// EAGLE variant the router doesn't know yet, which would otherwise silently
/// zero out cache-aware routing (the exact failure this whole path fixes).
/// Recognized non-EAGLE algorithms (and the absent field) map to `false`
/// silently.
pub(crate) fn classify_bigram(speculative_algorithm: Option<&str>, worker_url: &str) -> bool {
let Some(algo) = speculative_algorithm else {
return false;
};
let upper = algo.to_ascii_uppercase();
match upper.as_str() {
"EAGLE" | "EAGLE3" | "FROZEN_KV_MTP" => true,
_ => {
if upper.contains("EAGLE") || upper.contains("MTP") {
tracing::warn!(
worker_url = %worker_url,
speculative_algorithm = %algo,
"kv-events: unrecognized EAGLE-like speculative_algorithm; treating as \
non-bigram (unigram) hashing. If this is an EAGLE-family algorithm, \
cache-aware routing will silently never match — add it to classify_bigram",
);
}
false
}
}
} }
#[derive(Deserialize)] #[derive(Deserialize)]
@@ -270,10 +320,47 @@ mod tests {
topic: "kv".to_string(), topic: "kv".to_string(),
block_size: 64, block_size: 64,
dp_size: 2, dp_size: 2,
is_bigram: false,
}) })
); );
} }
/// EAGLE-family `speculative_algorithm` (and only those) must set
/// `is_bigram`, so the router selects the bigram hasher and its query
/// hashes match the worker's bigram-stored block hashes.
#[tokio::test]
async fn fetch_sets_is_bigram_for_eagle_family_only() {
for (algo, expected) in [
(Some("EAGLE"), true),
(Some("EAGLE3"), true),
(Some("FROZEN_KV_MTP"), true),
(Some("eagle"), true), // case-insensitive
(Some("NONE"), false),
(Some("NEXTN"), false), // non-eagle speculative algorithm
(None, false), // no speculative decoding
] {
let mut obj = json!({
"kv_events": {
"publisher": "zmq",
"endpoint_host": "*",
"endpoint_port_base": 5557,
"topic": "",
"block_size": 64,
"dp_size": 1,
}
});
if let Some(a) = algo {
obj["speculative_algorithm"] = json!(a);
}
let (url, _shutdown) = spawn_fake_worker(Arc::new(obj)).await;
let got = fetch_event_config(&url, &client()).await.unwrap().unwrap();
assert_eq!(
got.is_bigram, expected,
"speculative_algorithm={algo:?} should map to is_bigram={expected}"
);
}
}
/// Worker reports a specific bind host (not wildcard): gateway must /// Worker reports a specific bind host (not wildcard): gateway must
/// honour it instead of overwriting from the URL. /// honour it instead of overwriting from the URL.
#[tokio::test] #[tokio::test]
@@ -34,10 +34,11 @@
//! //!
//! ### Bigram mode //! ### Bigram mode
//! //!
//! Not supported in v1. SGLang's bigram mode interleaves overlapping //! EAGLE-family workers (`is_bigram = is_eagle`) hash KV blocks over
//! `(t_i, t_{i+1})` pairs into the hash. The gateway does not need this //! overlapping `(t_i, t_{i+1})` token pairs. That path is implemented as a
//! today; if/when it does, add a separate `compute_block_hashes_bigram` rather //! separate [`compute_block_hashes_bigram`] (below) rather than branching
//! than complicating the non-bigram fast path. //! inside the non-bigram fast path; `CacheAwareZmqPolicy::select` chooses
//! between the two from the worker-reported bigram flag.
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
@@ -107,10 +108,157 @@ pub fn sha256_to_i64(digest: &[u8; 32]) -> i64 {
i64::from_be_bytes(top) i64::from_be_bytes(top)
} }
/// Bigram variant of [`compute_block_hashes`], matching SGLang's `radix_cache`
/// worker when the model runs **EAGLE speculative decoding** (`is_bigram =
/// is_eagle`). Mirrors `RadixKey.hash_page` (Python:
/// `mem_cache/radix_cache.py`) on the bigram path:
///
/// - The logical sequence is the `N-1` overlapping bigrams of `N` raw tokens,
/// so the page count is `ceil((len-1) / block_size)`. Fewer than 2 tokens
/// yields no blocks.
/// - Each page `[start, end)` (in bigram-index space) feeds **both** tokens of
/// every bigram into the SHA256 hasher — `t[j]` then `t[j+1]`, each as 4
/// little-endian bytes — vs. the unigram path's single token per unit.
/// - Pages chain on the prior page's full 32-byte digest and truncate to i64
/// exactly as the unigram path does.
///
/// Use this (instead of [`compute_block_hashes`]) when the worker advertises an
/// EAGLE speculative algorithm via `/server_info`; otherwise the router's query
/// hashes won't match the worker's stored bigram block hashes and cache-aware
/// routing silently degrades to min-load.
pub fn compute_block_hashes_bigram(token_ids: &[u32], block_size: usize) -> Vec<i64> {
assert!(block_size > 0, "block_size must be positive");
// N raw tokens -> N-1 overlapping bigrams; fewer than 2 tokens -> no blocks.
let logical_len = token_ids.len().saturating_sub(1);
if logical_len == 0 {
return Vec::new();
}
let num_blocks = logical_len.div_ceil(block_size);
let mut out = Vec::with_capacity(num_blocks);
let mut prior: Option<[u8; 32]> = None;
let mut start = 0;
while start < logical_len {
let end = (start + block_size).min(logical_len);
let digest = chain_block_bigram(prior.as_ref(), token_ids, start, end);
out.push(sha256_to_i64(&digest));
prior = Some(digest);
start = end;
}
out
}
/// Hash a single bigram page: for each unit `j` in `[start, end)`, feed
/// `tokens[j]` then `tokens[j + 1]` (4 little-endian bytes each), chained on the
/// parent block's full 32-byte digest. The caller guarantees `end <= len - 1`,
/// so `tokens[j + 1]` is always in bounds. Mirrors the engine's `hash_page`
/// bigram branch.
#[inline]
fn chain_block_bigram(
parent_digest: Option<&[u8; 32]>,
tokens: &[u32],
start: usize,
end: usize,
) -> [u8; 32] {
let mut hasher = Sha256::new();
if let Some(parent) = parent_digest {
hasher.update(parent);
}
for j in start..end {
hasher.update(tokens[j].to_le_bytes());
hasher.update(tokens[j + 1].to_le_bytes());
}
hasher.finalize().into()
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
// ---- Bigram cross-language goldens ----
// Values produced by SGLang's REAL `RadixKey(..., is_bigram=True).hash_page`
// + `compute_node_hash_values` chunking + `hash_str_to_int64`, run against
// the deployed DeepSeek-V4-Flash engine. These lock byte-exact equivalence
// with the worker's stored block hashes — the contract that makes
// cache-aware routing actually match for EAGLE/bigram models.
#[test]
fn bigram_golden_single_block_full() {
// engine: chain_bigram([10,20,30,40], 4) -> [-2735951481331064195]
assert_eq!(
compute_block_hashes_bigram(&[10, 20, 30, 40], 4),
vec![-2735951481331064195_i64]
);
}
#[test]
fn bigram_golden_multi_block() {
// engine: chain_bigram([10,20,30,40,50], 2) -> [-8847804484166691499, 4989791362144317498]
assert_eq!(
compute_block_hashes_bigram(&[10, 20, 30, 40, 50], 2),
vec![-8847804484166691499_i64, 4989791362144317498_i64]
);
}
#[test]
fn bigram_golden_partial_last_block() {
// engine: chain_bigram([1,2,3,4,5,6], 4) -> [-638950109823820341, 3604587133525381017]
assert_eq!(
compute_block_hashes_bigram(&[1, 2, 3, 4, 5, 6], 4),
vec![-638950109823820341_i64, 3604587133525381017_i64]
);
}
#[test]
fn bigram_golden_longer_multi_block() {
// engine: chain_bigram([5,6,7,8,9,10,11,12,13], 4) -> [-2900568514773989563, -322435596280658912]
assert_eq!(
compute_block_hashes_bigram(&[5, 6, 7, 8, 9, 10, 11, 12, 13], 4),
vec![-2900568514773989563_i64, -322435596280658912_i64]
);
}
#[test]
fn bigram_single_bigram_equals_unigram_pair() {
// One bigram (10,20) feeds bytes 10,20 — identical to a unigram block
// [10,20]. engine: chain_bigram([10,20], 4) -> [978178666101069530],
// which equals the unigram block hash of [10,20].
assert_eq!(
compute_block_hashes_bigram(&[10, 20], 4),
vec![978178666101069530_i64]
);
assert_eq!(
compute_block_hashes_bigram(&[10, 20], 4),
compute_block_hashes(&[10, 20], 4)
);
}
#[test]
fn bigram_fewer_than_two_tokens_yields_no_blocks() {
// N tokens -> N-1 bigrams; <2 tokens -> 0 bigrams -> empty.
assert!(compute_block_hashes_bigram(&[10], 4).is_empty());
assert!(compute_block_hashes_bigram(&[], 4).is_empty());
}
#[test]
fn bigram_differs_from_unigram_for_multi_token_blocks() {
// Sanity: for >2 tokens the bigram hash must NOT equal the unigram hash
// (different byte stream) — this is exactly why a unigram-hashing
// router gets zero overlap against a bigram worker.
let toks = [10u32, 20, 30, 40];
assert_ne!(
compute_block_hashes_bigram(&toks, 4),
compute_block_hashes(&toks, 4)
);
}
#[test]
#[should_panic(expected = "block_size must be positive")]
fn bigram_zero_block_size_panics() {
let _ = compute_block_hashes_bigram(&[1, 2, 3], 0);
}
/// Helper for tests: derive the expected i64 from a list of tokens /// Helper for tests: derive the expected i64 from a list of tokens
/// chained against an optional parent digest. This mirrors `chain_block` /// chained against an optional parent digest. This mirrors `chain_block`
/// but is duplicated here so a regression in the production helper /// but is duplicated here so a regression in the production helper
@@ -204,11 +204,16 @@ impl KvEventIndex {
); );
return; return;
} }
// Establish the bigram flag alongside block_size. EAGLE-family workers
// hash KV blocks over token bigrams, so the policy must use the bigram
// hasher for its query hashes to match the worker's stored hashes.
self.block_size_oracle.set_bigram(cfg.is_bigram);
info!( info!(
worker_url = %worker_url, worker_url = %worker_url,
dp_size = cfg.dp_size, dp_size = cfg.dp_size,
port_base = cfg.port_base, port_base = cfg.port_base,
block_size = cfg.block_size, block_size = cfg.block_size,
is_bigram = cfg.is_bigram,
"kv-events: subscribing", "kv-events: subscribing",
); );
// Compute the DP ranks that will actually be subscribed (skip // Compute the DP ranks that will actually be subscribed (skip
@@ -646,6 +651,7 @@ mod tests {
topic: String::new(), topic: String::new(),
block_size: 128, block_size: 128,
dp_size: 1, dp_size: 1,
is_bigram: false,
}; };
index index
.add_worker("http://127.0.0.1:30100", Some(bad_cfg)) .add_worker("http://127.0.0.1:30100", Some(bad_cfg))
@@ -674,9 +680,34 @@ mod tests {
topic: String::new(), topic: String::new(),
block_size: 64, block_size: 64,
dp_size: 0, dp_size: 0,
is_bigram: false,
}; };
index.add_worker("http://127.0.0.1:30200", Some(cfg)).await; index.add_worker("http://127.0.0.1:30200", Some(cfg)).await;
assert_eq!(index.block_size_oracle().get(), Some(64)); assert_eq!(index.block_size_oracle().get(), Some(64));
index.shutdown().await; index.shutdown().await;
} }
#[tokio::test]
async fn add_worker_seeds_bigram_flag_from_event_config() {
// The discovery->routing seam: add_worker must publish
// EventConfig.is_bigram into the oracle (alongside block_size) so
// select() picks the bigram hasher for EAGLE workers.
let index = KvEventIndex::new();
assert!(!index.block_size_oracle().is_bigram());
// dp_size=0 short-circuits the subscriber spawn but still runs the seed.
let cfg = EventConfig {
host: "127.0.0.1".into(),
port_base: 30300,
topic: String::new(),
block_size: 64,
dp_size: 0,
is_bigram: true,
};
index.add_worker("http://127.0.0.1:30300", Some(cfg)).await;
assert!(
index.block_size_oracle().is_bigram(),
"add_worker must seed the bigram flag from EventConfig"
);
index.shutdown().await;
}
} }
@@ -23,8 +23,9 @@ pub mod tree;
pub mod wire; pub mod wire;
pub use block_size_oracle::BlockSizeOracle; pub use block_size_oracle::BlockSizeOracle;
pub(crate) use discovery::classify_bigram;
pub use discovery::{fetch_event_config, EventConfig}; pub use discovery::{fetch_event_config, EventConfig};
pub use hash::{compute_block_hashes, sha256_to_i64}; pub use hash::{compute_block_hashes, compute_block_hashes_bigram, sha256_to_i64};
pub use index::KvEventIndex; pub use index::KvEventIndex;
pub use subscriber::{KvEventSubscriberRegistry, WorkerEvent}; pub use subscriber::{KvEventSubscriberRegistry, WorkerEvent};
pub use tree::{HashTree, KvWorkerId, MatchResult}; pub use tree::{HashTree, KvWorkerId, MatchResult};
@@ -593,6 +593,7 @@ mod tests {
topic: String::new(), topic: String::new(),
block_size: 64, block_size: 64,
dp_size, dp_size,
is_bigram: false,
} }
} }
@@ -213,7 +213,64 @@ impl<'de> Deserialize<'de> for BoundedI64Vec {
} }
} }
/// `BoundedI64Vec`'s `u32` twin. Same shape, different cap. /// One element of a `token_ids` array. SGLang emits a flat `u32` per token for
/// unigram pages, but a 2-element `[t_i, t_{i+1}]` array per token for *bigram*
/// pages (`mem_cache/events.py`, `is_bigram` branch — DeepSeek-V4-class models).
/// `token_ids` is purely informational for the gateway (routing keys off the
/// engine-provided `block_hashes`), so we accept either shape and flatten the
/// ints rather than model the bigram pairing.
enum TokenCell {
One(u32),
Many(Vec<u32>),
}
impl<'de> Deserialize<'de> for TokenCell {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct V;
impl<'de> Visitor<'de> for V {
type Value = TokenCell;
fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str("a token id (u32) or an array of token ids")
}
// serde's default visit_u8/u16/u32 forward to visit_u64, and
// visit_i8/i16/i32 forward to visit_i64, so these two cover every
// integer width msgpack might use for a scalar token id.
fn visit_u64<E: de::Error>(self, v: u64) -> Result<TokenCell, E> {
Ok(TokenCell::One(v as u32))
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<TokenCell, E> {
Ok(TokenCell::One(v as u32))
}
fn visit_seq<A>(self, mut seq: A) -> Result<TokenCell, A::Error>
where
A: SeqAccess<'de>,
{
let mut ts: Vec<u32> = match seq.size_hint() {
Some(h) => Vec::with_capacity(h.min(8)),
None => Vec::new(),
};
while let Some(t) = seq.next_element::<u32>()? {
if ts.len() >= MAX_TOKENS_PER_EVENT {
return Err(de::Error::custom(format!(
"{PAYLOAD_TOO_LARGE_TAG}:token_ids:{}:{MAX_TOKENS_PER_EVENT}",
ts.len() + 1
)));
}
ts.push(t);
}
Ok(TokenCell::Many(ts))
}
}
deserializer.deserialize_any(V)
}
}
/// `BoundedI64Vec`'s `u32` twin. Same shape, different cap. Accepts both flat
/// (unigram) token ids and bigram `[t_i, t_{i+1}]` pairs via [`TokenCell`],
/// flattening the latter.
#[derive(Debug, Clone, PartialEq)] #[derive(Debug, Clone, PartialEq)]
struct BoundedU32Vec(Vec<u32>); struct BoundedU32Vec(Vec<u32>);
@@ -243,14 +300,27 @@ impl<'de> Deserialize<'de> for BoundedU32Vec {
Some(h) => Vec::with_capacity(h), Some(h) => Vec::with_capacity(h),
None => Vec::new(), None => Vec::new(),
}; };
while let Some(v) = seq.next_element::<u32>()? { // Each element is either a scalar token id (unigram) or a
// `[t_i, t_{i+1}]` pair (bigram); flatten both into `out`.
while let Some(cell) = seq.next_element::<TokenCell>()? {
let push = |t: u32, out: &mut Vec<u32>| -> Result<(), A::Error> {
if out.len() >= MAX_TOKENS_PER_EVENT { if out.len() >= MAX_TOKENS_PER_EVENT {
return Err(de::Error::custom(format!( return Err(de::Error::custom(format!(
"{PAYLOAD_TOO_LARGE_TAG}:token_ids:{}:{MAX_TOKENS_PER_EVENT}", "{PAYLOAD_TOO_LARGE_TAG}:token_ids:{}:{MAX_TOKENS_PER_EVENT}",
out.len() + 1 out.len() + 1
))); )));
} }
out.push(v); out.push(t);
Ok(())
};
match cell {
TokenCell::One(t) => push(t, &mut out)?,
TokenCell::Many(ts) => {
for t in ts {
push(t, &mut out)?;
}
}
}
} }
Ok(out) Ok(out)
} }
@@ -417,6 +487,84 @@ mod tests {
} }
} }
/// Encode `token_ids` the way SGLang's *bigram* pages do: a sequence of
/// 2-element `[t_i, t_{i+1}]` arrays instead of flat ints. See
/// `mem_cache/events.py` (`is_bigram` branch).
fn write_bigram_token_array(buf: &mut Vec<u8>, pairs: &[(u32, u32)]) {
mp::write_array_len(buf, pairs.len() as u32).unwrap();
for (a, b) in pairs {
mp::write_array_len(buf, 2).unwrap();
mp::write_uint(buf, *a as u64).unwrap();
mp::write_uint(buf, *b as u64).unwrap();
}
}
/// Like `build_block_stored_bytes`, but `token_ids` is the bigram
/// list-of-pairs shape that DeepSeek-V4-class models emit.
fn build_block_stored_bigram_bytes(
block_hashes: &[i64],
parent: Option<i64>,
token_pairs: &[(u32, u32)],
block_size: u32,
lora_id: Option<i64>,
medium: Option<&str>,
) -> Vec<u8> {
let mut buf = Vec::new();
write_event_array(&mut buf, "BlockStored", 7);
write_i64_array(&mut buf, block_hashes);
match parent {
Some(v) => {
mp::write_sint(&mut buf, v).unwrap();
}
None => mp::write_nil(&mut buf).unwrap(),
}
write_bigram_token_array(&mut buf, token_pairs);
mp::write_uint(&mut buf, block_size as u64).unwrap();
match lora_id {
Some(v) => {
mp::write_sint(&mut buf, v).unwrap();
}
None => mp::write_nil(&mut buf).unwrap(),
}
match medium {
Some(s) => mp::write_str(&mut buf, s).unwrap(),
None => mp::write_nil(&mut buf).unwrap(),
}
buf
}
/// Regression: bigram models (e.g. DeepSeek-V4-Flash) emit `token_ids` as
/// `[[t_i, t_{i+1}], ...]`. The decoder previously read `token_ids` as a
/// flat `u32` array and failed the entire batch with
/// "wrong msgpack marker FixArray(2)", silently disabling cache-aware
/// routing. It must instead accept the bigram shape (flattening the ints).
#[test]
fn decodes_block_stored_with_bigram_token_ids() {
let event = build_block_stored_bigram_bytes(
&[111_i64],
None,
&[(10, 20), (20, 30)],
2,
None,
Some("GPU"),
);
let bytes = build_batch_bytes(1.5, &[event], Some(0), true);
let batch = decode_event_batch(&bytes).expect("decode bigram token_ids");
assert_eq!(batch.events.len(), 1);
match &batch.events[0] {
KvCacheEvent::BlockStored(b) => {
// routing-relevant fields decode unchanged
assert_eq!(b.block_hashes, vec![111]);
assert_eq!(b.parent_block_hash, None);
assert_eq!(b.block_size, 2);
// bigram pairs are flattened into the (informational) token vec
assert_eq!(b.token_ids, vec![10, 20, 20, 30]);
}
other => panic!("expected BlockStored, got {other:?}"),
}
}
/// Build a full BlockStored event as msgspec would emit it (all 7 /// Build a full BlockStored event as msgspec would emit it (all 7
/// elements: tag + 6 fields). `medium` may be Some/None. /// elements: tag + 6 fields). `medium` may be Some/None.
fn build_block_stored_bytes( fn build_block_stored_bytes(
@@ -11,6 +11,7 @@ pub mod registry;
pub mod round_robin; pub mod round_robin;
use crate::discovery::ModelId; use crate::discovery::ModelId;
use crate::server::metrics::MetricsRegistry;
use crate::workers::Worker; use crate::workers::Worker;
use dashmap::DashMap; use dashmap::DashMap;
use std::sync::Arc; use std::sync::Arc;
@@ -46,6 +47,13 @@ impl<'a> SelectionContext<'a> {
pub trait Policy: Send + Sync + std::fmt::Debug { pub trait Policy: Send + Sync + std::fmt::Debug {
fn select(&self, workers: &[Arc<Worker>], ctx: &SelectionContext<'_>) -> Option<Arc<Worker>>; fn select(&self, workers: &[Arc<Worker>], ctx: &SelectionContext<'_>) -> Option<Arc<Worker>>;
/// Attach the process metrics registry after construction. Default is a
/// no-op — only policies that emit metrics (cache-aware-zmq's
/// `sgl_router_overlap_blocks`) override it. Mirrors
/// `ActiveLoadRegistry::attach_metrics`: the registry is built after the
/// policies, so it is injected here rather than passed to the constructor.
fn attach_metrics(&self, _metrics: Arc<MetricsRegistry>) {}
} }
#[derive(Debug, Default)] #[derive(Debug, Default)]
@@ -61,4 +69,13 @@ impl PolicyRegistry {
pub fn get(&self, model: &ModelId) -> Option<Arc<dyn Policy>> { pub fn get(&self, model: &ModelId) -> Option<Arc<dyn Policy>> {
self.by_model.get(model).map(|p| p.clone()) self.by_model.get(model).map(|p| p.clone())
} }
/// Inject the metrics registry into every registered policy. Called once
/// at startup (after the registry is built) so metrics-emitting policies
/// can record into the shared registry.
pub fn attach_metrics(&self, metrics: Arc<MetricsRegistry>) {
for entry in self.by_model.iter() {
entry.value().attach_metrics(Arc::clone(&metrics));
}
}
} }
@@ -69,6 +69,11 @@ impl AppContext {
// Without this, the metric is permanently 0 in production even // Without this, the metric is permanently 0 in production even
// though the chat handler is faithfully calling `register`. // though the chat handler is faithfully calling `register`.
active_load.attach_metrics(Arc::clone(&metrics)); 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.
policies.attach_metrics(Arc::clone(&metrics));
Self { Self {
config, config,
tokenizers, tokenizers,
@@ -100,14 +105,18 @@ impl AppContext {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![], model: crate::config::ModelConfig {
discovery: crate::config::DiscoveryConfig { id: "stub-model".into(),
backend: crate::config::DiscoveryBackend::StaticUrls( tokenizer_path: "stub".into(),
policy: crate::config::PolicyKind::RoundRobin,
circuit_breaker: None,
cache_aware: None,
},
discovery: crate::config::DiscoveryBackend::StaticUrls(
crate::config::StaticUrlsDiscoveryConfig { crate::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}, },
), ),
},
proxy: crate::config::ProxyConfig::default(), proxy: crate::config::ProxyConfig::default(),
active_load: crate::config::ActiveLoadConfig::default(), active_load: crate::config::ActiveLoadConfig::default(),
}, },
@@ -144,6 +144,13 @@ impl ApiError {
ApiError::Internal(_) => (StatusCode::INTERNAL_SERVER_ERROR, "internal_error"), ApiError::Internal(_) => (StatusCode::INTERNAL_SERVER_ERROR, "internal_error"),
} }
} }
/// The HTTP status this error maps to — same value the client receives via
/// `into_response`. Exposed so the access log records the real status
/// (e.g. 502/503/504) instead of a sentinel.
pub fn status_code(&self) -> StatusCode {
self.status_and_code().0
}
} }
#[derive(Serialize)] #[derive(Serialize)]
@@ -75,6 +75,7 @@ pub async fn chat_completions(
headers: HeaderMap, headers: HeaderMap,
body: Bytes, body: Bytes,
) -> Result<Response<Body>, ApiError> { ) -> Result<Response<Body>, ApiError> {
let start = std::time::Instant::now();
let probe = parse_probe(&body)?; let probe = parse_probe(&body)?;
let streaming = probe.stream.unwrap_or(false); let streaming = probe.stream.unwrap_or(false);
let model_str = probe let model_str = probe
@@ -397,6 +398,37 @@ pub async fn chat_completions(
ctx.metrics ctx.metrics
.record_request(&metrics_worker_url, &metrics_model, metrics_mode, outcome); .record_request(&metrics_worker_url, &metrics_model, metrics_mode, outcome);
// Per-request access log — always on at INFO so incoming traffic and its
// status are visible without DEBUG. `request_id` is the client/gateway
// X-Request-Id (echoed end-to-end); `worker` is the engine the policy
// selected. The cache-aware routing rationale is logged separately at
// DEBUG by the policy.
let request_id = headers
.get("x-request-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("-");
let http_status = match &result {
Ok(resp) => resp.status().as_u16(),
Err(e) => e.status_code().as_u16(),
};
let outcome_str = match outcome {
RequestOutcome::Success => "success",
RequestOutcome::Error => "error",
RequestOutcome::Cancelled => "cancelled",
};
tracing::info!(
request_id = %request_id,
method = "POST",
path = "/v1/chat/completions",
model = %metrics_model,
worker = %metrics_worker_url,
outcome = outcome_str,
http_status,
stream = streaming,
latency_ms = start.elapsed().as_millis() as u64,
"chat_completions",
);
// Mirror the upstream `x-sgl-decode-url` hint onto the response so // 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 PD decode affinity without
// sniffing the proxy hop. The request-side header was set above for // sniffing the proxy hop. The request-side header was set above for
@@ -21,16 +21,14 @@ pub struct ModelEntry {
} }
pub async fn list_models(State(ctx): State<Arc<AppContext>>) -> Json<ModelsList> { pub async fn list_models(State(ctx): State<Arc<AppContext>>) -> Json<ModelsList> {
let data = ctx // The router serves a single configured model; OpenAI clients still
.config // expect a list shape, so return a one-element `data` array.
.models let m = &ctx.config.model;
.iter() let data = vec![ModelEntry {
.map(|m| ModelEntry {
id: m.id.clone(), id: m.id.clone(),
object: "model", object: "model",
owned_by: "sglang", owned_by: "sglang",
}) }];
.collect();
Json(ModelsList { Json(ModelsList {
object: "list", object: "list",
data, data,
@@ -47,24 +45,15 @@ mod tests {
use crate::config::PolicyKind; use crate::config::PolicyKind;
#[tokio::test] #[tokio::test]
async fn lists_configured_models() { async fn lists_configured_model() {
let mut ctx = crate::server::app_context::AppContext::stub(); let mut ctx = crate::server::app_context::AppContext::stub();
ctx.config.models = vec![ ctx.config.model = crate::config::ModelConfig {
crate::config::ModelConfig {
id: "qwen3".into(), id: "qwen3".into(),
tokenizer_path: "x".into(), tokenizer_path: "x".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}, };
crate::config::ModelConfig {
id: "deepseek".into(),
tokenizer_path: "y".into(),
policy: PolicyKind::RoundRobin,
circuit_breaker: None,
cache_aware: None,
},
];
let app = crate::server::app::build_router(std::sync::Arc::new(ctx)); let app = crate::server::app::build_router(std::sync::Arc::new(ctx));
let res = app let res = app
.oneshot( .oneshot(
@@ -85,13 +74,12 @@ mod tests {
.iter() .iter()
.map(|m| m["id"].as_str().unwrap()) .map(|m| m["id"].as_str().unwrap())
.collect(); .collect();
assert_eq!(ids, vec!["qwen3", "deepseek"]); assert_eq!(ids, vec!["qwen3"]);
assert_eq!(v["data"][0]["object"], "model"); assert_eq!(v["data"][0]["object"], "model");
// Pin `owned_by` so a refactor that flips the hardcoded value to // Pin `owned_by` so a refactor that flips the hardcoded value to
// "openai" / "" / a typo would fail loudly here. OpenAI clients // "openai" / "" / a typo would fail loudly here. OpenAI clients
// expect this field and some (e.g. langchain-openai) treat // expect this field and some (e.g. langchain-openai) treat
// `owned_by != "system"` as a meaningful signal. // `owned_by != "system"` as a meaningful signal.
assert_eq!(v["data"][0]["owned_by"], "sglang"); assert_eq!(v["data"][0]["owned_by"], "sglang");
assert_eq!(v["data"][1]["owned_by"], "sglang");
} }
} }
@@ -114,20 +114,18 @@ mod tests {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![crate::config::ModelConfig { model: crate::config::ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: crate::config::DiscoveryConfig { discovery: crate::config::DiscoveryBackend::StaticUrls(
backend: crate::config::DiscoveryBackend::StaticUrls(
crate::config::StaticUrlsDiscoveryConfig { crate::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}, },
), ),
},
proxy: crate::config::ProxyConfig::default(), proxy: crate::config::ProxyConfig::default(),
active_load: crate::config::ActiveLoadConfig::default(), active_load: crate::config::ActiveLoadConfig::default(),
}; };
@@ -3,12 +3,64 @@
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use dynamo_tokenizers::{traits::DecodeResult, Tokenizer}; use dynamo_tokenizers::{traits::DecodeResult, Tokenizer};
use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
pub fn load(path: &str) -> Result<Arc<Tokenizer>> { /// Load a tokenizer from `source`, which is either a local `tokenizer.json`
/// path or a HuggingFace repo id.
///
/// An existing local file (or anything with a filesystem-path shape) is
/// loaded directly via `Tokenizer::from_file`. Otherwise `source` is treated
/// as a HuggingFace repo id and its `tokenizer.json` is downloaded (once, at
/// startup) into the HF cache, honoring `HF_TOKEN` / `HF_HOME` /
/// `HF_HUB_OFFLINE`. `dynamo_tokenizers` itself has no HF-download path, so
/// the fetch is done here via `hf-hub`.
pub fn load(source: &str) -> Result<Arc<Tokenizer>> {
if Path::new(source).is_file() || looks_like_path(source) {
return Tokenizer::from_file(source)
.map(Arc::new)
.with_context(|| format!("load tokenizer from {source}"));
}
let downloaded = download_tokenizer_json(source)?;
let path = downloaded
.to_str()
.context("downloaded tokenizer path is not valid UTF-8")?;
Tokenizer::from_file(path) Tokenizer::from_file(path)
.map(Arc::new) .map(Arc::new)
.with_context(|| format!("load tokenizer from {path}")) .with_context(|| format!("load downloaded tokenizer for {source}"))
}
/// Treat `source` as a filesystem path (rather than a HuggingFace repo id)
/// when it has a path-like shape — an absolute/relative prefix or a `.json`
/// suffix. HF repo ids are `namespace/name` with none of these markers, so a
/// missing local file like `/models/tok.json` reports a load error instead of
/// silently attempting a (doomed) network fetch.
fn looks_like_path(source: &str) -> bool {
source.starts_with('/')
|| source.starts_with("./")
|| source.starts_with("../")
|| source.starts_with('~')
|| source.ends_with(".json")
}
/// Download `tokenizer.json` for a HuggingFace repo id and return the cached
/// local path. Uses the blocking `ureq` API (this runs once at startup,
/// before the server begins serving) and `from_env` so `HF_TOKEN` /
/// `HF_HOME` / endpoint overrides are honored.
fn download_tokenizer_json(repo_id: &str) -> Result<std::path::PathBuf> {
use hf_hub::api::sync::ApiBuilder;
let api = ApiBuilder::from_env()
.build()
.context("initialize HuggingFace Hub client")?;
api.model(repo_id.to_string())
.get("tokenizer.json")
.with_context(|| {
format!(
"download tokenizer.json for HuggingFace repo {repo_id:?} \
(pass --tokenizer-path with a local tokenizer.json, or set HF_TOKEN \
for a gated/private repo)"
)
})
} }
pub fn encode(t: &Tokenizer, text: &str) -> Result<Vec<u32>> { pub fn encode(t: &Tokenizer, text: &str) -> Result<Vec<u32>> {
+5 -8
View File
@@ -24,10 +24,9 @@ impl std::fmt::Debug for TokenizerRegistry {
impl TokenizerRegistry { impl TokenizerRegistry {
pub fn load_from_config(cfg: &crate::config::Config) -> Result<Self> { pub fn load_from_config(cfg: &crate::config::Config) -> Result<Self> {
let me = TokenizerRegistry::default(); let me = TokenizerRegistry::default();
for m in &cfg.models { let m = &cfg.model;
let t = adapter::load(&m.tokenizer_path)?; let t = adapter::load(&m.tokenizer_path)?;
me.inner.insert(m.id.clone(), t); me.inner.insert(m.id.clone(), t);
}
Ok(me) Ok(me)
} }
@@ -54,20 +53,18 @@ mod tests {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![crate::config::ModelConfig { model: crate::config::ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: crate::config::DiscoveryConfig { discovery: crate::config::DiscoveryBackend::StaticUrls(
backend: crate::config::DiscoveryBackend::StaticUrls(
crate::config::StaticUrlsDiscoveryConfig { crate::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}, },
), ),
},
proxy: crate::config::ProxyConfig::default(), proxy: crate::config::ProxyConfig::default(),
active_load: crate::config::ActiveLoadConfig::default(), active_load: crate::config::ActiveLoadConfig::default(),
} }
@@ -192,7 +189,7 @@ mod tests {
#[test] #[test]
fn missing_file_errors() { fn missing_file_errors() {
let mut c = cfg(); let mut c = cfg();
c.models[0].tokenizer_path = "/nonexistent.json".into(); c.model.tokenizer_path = "/nonexistent.json".into();
let err = TokenizerRegistry::load_from_config(&c).unwrap_err(); let err = TokenizerRegistry::load_from_config(&c).unwrap_err();
assert!(err.to_string().to_lowercase().contains("tokenizer")); assert!(err.to_string().to_lowercase().contains("tokenizer"));
} }
@@ -125,9 +125,15 @@ impl WorkerIntrospector {
None => None, None => None,
}; };
// EAGLE-family speculative decoding ⇒ the worker hashes KV blocks over
// token bigrams; the router must mirror that on the selection side.
let is_bigram = crate::policies::kv_events::classify_bigram(
parsed.speculative_algorithm.as_deref(),
worker_url,
);
let event_config = parsed let event_config = parsed
.kv_events .kv_events
.map(|block| resolve_event_config(block, worker_url)); .map(|block| resolve_event_config(block, worker_url, is_bigram));
let disaggregation_role = resolve_disaggregation_role( let disaggregation_role = resolve_disaggregation_role(
parsed.disaggregation_mode.as_deref(), parsed.disaggregation_mode.as_deref(),
@@ -265,7 +271,11 @@ impl Default for WorkerIntrospector {
/// unchanged: the subsequent ZMQ connect will fail visibly with the /// unchanged: the subsequent ZMQ connect will fail visibly with the
/// wildcard literal, which is the same observable failure mode that /// wildcard literal, which is the same observable failure mode that
/// would occur today if the bind/connect were skipped. /// would occur today if the bind/connect were skipped.
pub(crate) fn resolve_event_config(block: KvEventsBlock, worker_url: &str) -> EventConfig { pub(crate) fn resolve_event_config(
block: KvEventsBlock,
worker_url: &str,
is_bigram: bool,
) -> EventConfig {
let host = if matches!( let host = if matches!(
block.endpoint_host.as_str(), block.endpoint_host.as_str(),
"*" | "0.0.0.0" | "::" | "[::]" "*" | "0.0.0.0" | "::" | "[::]"
@@ -292,6 +302,7 @@ pub(crate) fn resolve_event_config(block: KvEventsBlock, worker_url: &str) -> Ev
topic: block.topic, topic: block.topic,
block_size: block.block_size, block_size: block.block_size,
dp_size: block.dp_size, dp_size: block.dp_size,
is_bigram,
} }
} }
@@ -304,6 +315,11 @@ struct ServerInfoBody {
served_model_name: Option<String>, served_model_name: Option<String>,
#[serde(default)] #[serde(default)]
kv_events: Option<KvEventsBlock>, kv_events: Option<KvEventsBlock>,
/// Top-level `speculative_algorithm`. EAGLE-family values
/// (EAGLE / EAGLE3 / FROZEN_KV_MTP) ⇒ the worker hashes KV blocks over
/// token bigrams. Absent on workers without speculative decoding.
#[serde(default)]
speculative_algorithm: Option<String>,
/// Carries the value of `ServerArgs.disaggregation_mode` /// Carries the value of `ServerArgs.disaggregation_mode`
/// (`"null"` | `"prefill"` | `"decode"`). Absent on older SGLang /// (`"null"` | `"prefill"` | `"decode"`). Absent on older SGLang
/// versions that predate the field. /// versions that predate the field.
@@ -369,6 +385,59 @@ mod tests {
WorkerIntrospector::new(Duration::from_millis(500)) WorkerIntrospector::new(Duration::from_millis(500))
} }
/// The PRIMARY `/server_info` path (the introspector, not the discovery.rs
/// fallback) must flag `is_bigram` for an EAGLE worker so the policy picks
/// the bigram hasher. Regression guard for the duplicated parse + the
/// `resolve_event_config(.., is_bigram)` threading.
#[tokio::test]
async fn fetch_sets_is_bigram_for_eagle_worker() {
let (url, _shutdown) = spawn_fake_worker(json!({
"served_model_name": "m",
"speculative_algorithm": "EAGLE",
"kv_events": {
"publisher": "zmq",
"endpoint_host": "*",
"endpoint_port_base": 5557,
"topic": "",
"block_size": 64,
"dp_size": 1,
}
}))
.await;
let cfg = fast_introspector()
.fetch(&url)
.await
.event_config
.expect("kv_events present");
assert!(
cfg.is_bigram,
"EAGLE worker via the introspector must set is_bigram"
);
}
/// A non-speculative worker (no `speculative_algorithm`) must NOT be bigram.
#[tokio::test]
async fn fetch_no_bigram_without_speculative_algorithm() {
let (url, _shutdown) = spawn_fake_worker(json!({
"served_model_name": "m",
"kv_events": {
"publisher": "zmq",
"endpoint_host": "*",
"endpoint_port_base": 5557,
"topic": "",
"block_size": 64,
"dp_size": 1,
}
}))
.await;
let cfg = fast_introspector()
.fetch(&url)
.await
.event_config
.expect("kv_events present");
assert!(!cfg.is_bigram, "non-speculative worker must not be bigram");
}
#[tokio::test] #[tokio::test]
async fn fetch_returns_both_served_model_name_and_event_config() { async fn fetch_returns_both_served_model_name_and_event_config() {
let (url, _shutdown) = spawn_fake_worker(json!({ let (url, _shutdown) = spawn_fake_worker(json!({
+11 -14
View File
@@ -16,19 +16,18 @@ use tokio::task::JoinHandle;
/// Resolve the circuit-breaker config for all model IDs carried by a spec. /// Resolve the circuit-breaker config for all model IDs carried by a spec.
/// ///
/// Workers may serve multiple models; we use the config of the **first** model /// The router serves a single configured model; apply its circuit-breaker
/// that has an explicit CB config, falling back to `None` (default config). /// config when this worker advertises that model id. Falls back to `None`
/// (default config) otherwise.
fn cb_config_for_spec(spec: &WorkerSpec, cfg: &Config) -> Option<CircuitBreakerConfig> { fn cb_config_for_spec(spec: &WorkerSpec, cfg: &Config) -> Option<CircuitBreakerConfig> {
for model_id in &spec.model_ids { let model = &cfg.model;
if let Some(mc) = cfg.models.iter().find(|m| m.id == model_id.0) { let cbc = model.circuit_breaker.as_ref()?;
if let Some(cbc) = &mc.circuit_breaker { if spec.model_ids.iter().any(|id| id.0 == model.id) {
return Some(CircuitBreakerConfig { return Some(CircuitBreakerConfig {
threshold: cbc.threshold, threshold: cbc.threshold,
cool_down: Duration::from_secs(cbc.cool_down_secs), cool_down: Duration::from_secs(cbc.cool_down_secs),
}); });
} }
}
}
None None
} }
@@ -279,8 +278,8 @@ async fn register_one(
mod tests { mod tests {
use super::*; use super::*;
use crate::config::{ use crate::config::{
ActiveLoadConfig, CircuitBreakerConfig as RawCbConfig, DiscoveryBackend, DiscoveryConfig, ActiveLoadConfig, CircuitBreakerConfig as RawCbConfig, DiscoveryBackend, ModelConfig,
ModelConfig, PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use crate::discovery::{WorkerId, WorkerMode}; use crate::discovery::{WorkerId, WorkerMode};
use axum::{routing::get, Json, Router}; use axum::{routing::get, Json, Router};
@@ -296,7 +295,7 @@ mod tests {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![ModelConfig { model: ModelConfig {
id: id.into(), id: id.into(),
tokenizer_path: "/tmp/x".into(), tokenizer_path: "/tmp/x".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
@@ -305,12 +304,10 @@ mod tests {
cool_down_secs, cool_down_secs,
}), }),
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://test:30000".into()], urls: vec!["http://test:30000".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }
@@ -90,8 +90,7 @@ async fn static_urls_pd_role_resolved_end_to_end() {
use axum::{routing::get, Json, Router}; use axum::{routing::get, Json, Router};
use serde_json::json; use serde_json::json;
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ObservabilityConfig, ProxyConfig, ServerConfig,
ProxyConfig, ServerConfig,
}; };
use sgl_router::discovery::{spawn_discovery, WorkerId}; use sgl_router::discovery::{spawn_discovery, WorkerId};
use sgl_router::workers::{manager, WorkerRegistry}; use sgl_router::workers::{manager, WorkerRegistry};
@@ -127,12 +126,16 @@ async fn static_urls_pd_role_resolved_end_to_end() {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![], model: sgl_router::config::ModelConfig {
discovery: DiscoveryConfig { id: "tiny".into(),
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig { tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: sgl_router::config::PolicyKind::RoundRobin,
circuit_breaker: None,
cache_aware: None,
},
discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec![url.clone()], urls: vec![url.clone()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
}; };
@@ -65,20 +65,18 @@ async fn zmq_indexer_routes_to_publishing_worker_e2e() {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![sgl_router::config::ModelConfig { model: sgl_router::config::ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: sgl_router::config::PolicyKind::CacheAwareZmq, policy: sgl_router::config::PolicyKind::CacheAwareZmq,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: sgl_router::config::DiscoveryConfig { discovery: sgl_router::config::DiscoveryBackend::StaticUrls(
backend: sgl_router::config::DiscoveryBackend::StaticUrls(
sgl_router::config::StaticUrlsDiscoveryConfig { sgl_router::config::StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}, },
), ),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
}; };
@@ -127,6 +125,7 @@ async fn zmq_indexer_routes_to_publishing_worker_e2e() {
topic: String::new(), topic: String::new(),
block_size, block_size,
dp_size: 1, dp_size: 1,
is_bigram: false,
}; };
kv_index.add_worker(url_a, Some(preresolved.clone())).await; kv_index.add_worker(url_a, Some(preresolved.clone())).await;
kv_index.add_worker(url_b, Some(preresolved)).await; kv_index.add_worker(url_b, Some(preresolved)).await;
@@ -41,6 +41,7 @@ async fn two_independent_subscribers_converge_to_same_tree_state() {
topic: String::new(), topic: String::new(),
block_size, block_size,
dp_size: 1, dp_size: 1,
is_bigram: false,
}; };
// 2. Two independent router-process surrogates, each with its own // 2. Two independent router-process surrogates, each with its own
@@ -173,6 +174,7 @@ async fn two_subscribers_merge_events_from_two_publishers() {
topic: String::new(), topic: String::new(),
block_size, block_size,
dp_size: 1, dp_size: 1,
is_bigram: false,
}; };
let cfg_y = EventConfig { let cfg_y = EventConfig {
host: "127.0.0.1".into(), host: "127.0.0.1".into(),
@@ -180,6 +182,7 @@ async fn two_subscribers_merge_events_from_two_publishers() {
topic: String::new(), topic: String::new(),
block_size, block_size,
dp_size: 1, dp_size: 1,
is_bigram: false,
}; };
// Both routers subscribe to BOTH workers — the production fan-out. // Both routers subscribe to BOTH workers — the production fan-out.
+23 -40
View File
@@ -166,81 +166,64 @@ def _find_tokenizer_path(model: str) -> str:
return model return model
def build_smoke_router_config( def build_smoke_router_args(
*, *,
host: str, host: str,
port: int, port: int,
model: str, model: str,
tokenizer_path: str, tokenizer_path: str,
sglang_url: str, sglang_url: str,
) -> str: ) -> list[str]:
"""Build the TOML the smoke `router` fixture writes to disk. """Build the sgl-router CLI flags the smoke ``router`` fixture launches.
Returns ``main_config_text`` carrying ``[server]``, ``[[models]]``, Static single-worker discovery (``--worker-urls``) pointed at the one
and ``[discovery] backend = "static_urls"`` with the worker URL SGLang worker, serving exactly one model.
inline. The Rust ``Config`` struct requires a ``[discovery]``
section (``DiscoveryConfig`` has no ``#[serde(default)]``) and has
no top-level ``workers`` field. The previous ``static_file``
backend was replaced by ``static_urls`` (which holds the URL list
inline rather than via a side-car file).
"""
return f"""\
[server]
host = "{host}"
port = {port}
[[models]]
id = "{model}"
tokenizer_path = "{tokenizer_path}"
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = ["{sglang_url}"]
""" """
return [
"--host",
host,
"--port",
str(port),
"--model-id",
model,
"--tokenizer-path",
tokenizer_path,
"--worker-urls",
sglang_url,
]
@pytest.fixture(scope="session") @pytest.fixture(scope="session")
def router(sglang_server): # noqa: ARG001 (sglang_server must start first) def router(sglang_server): # noqa: ARG001 (sglang_server must start first)
"""Launch sgl-router on port 8090 pointed at the SGLang worker.""" """Launch sgl-router on port 8090 pointed at the SGLang worker."""
tok_path = _find_tokenizer_path(MODEL) tok_path = _find_tokenizer_path(MODEL)
cfg_handle = tempfile.NamedTemporaryFile(mode="w", suffix=".toml", delete=False) args = build_smoke_router_args(
cfg_path = Path(cfg_handle.name)
main_text = build_smoke_router_config(
host="0.0.0.0", host="0.0.0.0",
port=ROUTER_PORT, port=ROUTER_PORT,
model=MODEL, model=MODEL,
tokenizer_path=tok_path, tokenizer_path=tok_path,
sglang_url=f"http://localhost:{SGLANG_PORT}", sglang_url=f"http://localhost:{SGLANG_PORT}",
) )
cfg_handle.write(main_text)
cfg_handle.close()
try:
proc = subprocess.Popen( proc = subprocess.Popen(
[str(_BINARY), "--config", str(cfg_path)], [str(_BINARY), *args],
stdout=subprocess.PIPE, stdout=subprocess.PIPE,
stderr=subprocess.STDOUT, stderr=subprocess.STDOUT,
) )
# try/finally so the router is always reaped — on a readiness-probe
# failure, a test-body error, or a session-teardown exception alike.
try: try:
_wait_http(f"http://localhost:{ROUTER_PORT}/readyz", timeout=60) _wait_http(f"http://localhost:{ROUTER_PORT}/readyz", timeout=60)
except Exception:
proc.send_signal(signal.SIGTERM)
proc.wait(timeout=30)
raise
yield f"http://localhost:{ROUTER_PORT}" yield f"http://localhost:{ROUTER_PORT}"
finally:
if proc.poll() is None:
proc.send_signal(signal.SIGTERM) proc.send_signal(signal.SIGTERM)
try: try:
proc.wait(timeout=30) proc.wait(timeout=30)
except subprocess.TimeoutExpired: except subprocess.TimeoutExpired:
proc.kill() proc.kill()
proc.wait() proc.wait()
finally:
cfg_path.unlink(missing_ok=True)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -1,23 +1,22 @@
"""Minimal sgl-router Gateway class — adapted from SMG's e2e_test/infra/gateway.py. """Minimal sgl-router Gateway class for e2e tests.
Differences from SMG: sgl-router is a Rust binary
- SMG drives a Python launcher (`python3 -m sglang_router.launch_router`) (`experimental/sgl-router/target/release/sgl-router`) configured entirely
with worker URLs on the CLI. through CLI flags. This Gateway execs the binary with `--worker-urls <...>`
- sgl-router uses a Rust binary (`experimental/sgl-router/target/release/sgl-router`) (static discovery) plus the model + policy flags.
with a TOML config file. Worker discovery is config-file-based; this
Gateway writes a TOML to a tempfile and execs the binary with
`--config <tempfile>`.
Supported lifecycles: Supported lifecycles:
- Regular mode: one model, N worker URLs, single policy. - Regular mode: one model, N worker URLs, single policy.
- PD mode: one model, prefill_workers + decode_workers (lists of URLs), - PD mode: one model; prefill + decode URLs all go into one
discovery emits separate `WorkerMode::Prefill` / `WorkerMode::Decode` `--worker-urls` static list. Each worker is seeded as
entries. The router resolves PD pool isolation at request time. `WorkerMode::Plain` and its actual prefill/decode role + bootstrap
port are resolved from `/server_info` introspection, after which the
router isolates the PD pools at request time.
Use as a context manager: Use as a context manager:
with Gateway() as gw: with Gateway() as gw:
gw.start_regular(model_path="...", worker_urls=[...]) gw.start_regular(model_id="...", tokenizer_path="...", worker_urls=[...])
resp = httpx.post(f"{gw.base_url}/v1/chat/completions", json=...) resp = httpx.post(f"{gw.base_url}/v1/chat/completions", json=...)
or pytest fixture style (see e2e_test/conftest.py). or pytest fixture style (see e2e_test/conftest.py).
@@ -30,7 +29,6 @@ import os
import signal import signal
import socket import socket
import subprocess import subprocess
import tempfile
import time import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
@@ -93,8 +91,11 @@ def _resolve_tokenizer_path(tokenizer_path: str) -> str:
cached = try_to_load_from_cache(tokenizer_path, "tokenizer.json") cached = try_to_load_from_cache(tokenizer_path, "tokenizer.json")
if cached and Path(cached).is_file(): if cached and Path(cached).is_file():
return str(cached) return str(cached)
except Exception: # noqa: BLE001 except Exception as exc: # noqa: BLE001
pass # A cache miss is normal; log other failures (corrupt cache,
# signature change) so a later tokenizer-load error is traceable
# rather than mysterious.
logger.debug("HF tokenizer cache lookup failed for %r: %s", tokenizer_path, exc)
return tokenizer_path return tokenizer_path
@@ -150,7 +151,6 @@ class Gateway:
self.stale_request_timeout_secs = stale_request_timeout_secs self.stale_request_timeout_secs = stale_request_timeout_secs
self.process: subprocess.Popen | None = None self.process: subprocess.Popen | None = None
self._config_path: Path | None = None
self._started: bool = False self._started: bool = False
# Track child workers we spawned so __exit__ can tear them down. # Track child workers we spawned so __exit__ can tear them down.
self._owned_workers: list[subprocess.Popen] = [] self._owned_workers: list[subprocess.Popen] = []
@@ -172,7 +172,6 @@ class Gateway:
tokenizer_path: str, tokenizer_path: str,
worker_urls: list[str], worker_urls: list[str],
policy: str = "round_robin", policy: str = "round_robin",
extra_models: list[dict] | None = None,
timeout: float = 60.0, timeout: float = 60.0,
) -> None: ) -> None:
"""Start the router in regular (non-PD) mode. """Start the router in regular (non-PD) mode.
@@ -190,12 +189,11 @@ class Gateway:
timeout: How long to wait for ``/readyz`` before giving up. timeout: How long to wait for ``/readyz`` before giving up.
""" """
self._launch( self._launch(
self._build_config( self._build_args(
model_id=model_id, model_id=model_id,
tokenizer_path=tokenizer_path, tokenizer_path=tokenizer_path,
urls=list(worker_urls), urls=list(worker_urls),
policy=policy, policy=policy,
extra_models=extra_models or [],
), ),
timeout=timeout, timeout=timeout,
) )
@@ -222,12 +220,11 @@ class Gateway:
assumed. assumed.
""" """
self._launch( self._launch(
self._build_config( self._build_args(
model_id=model_id, model_id=model_id,
tokenizer_path=tokenizer_path, tokenizer_path=tokenizer_path,
urls=list(prefill_urls) + list(decode_urls), urls=list(prefill_urls) + list(decode_urls),
policy=policy, policy=policy,
extra_models=[],
), ),
timeout=timeout, timeout=timeout,
) )
@@ -247,9 +244,6 @@ class Gateway:
except ProcessLookupError: except ProcessLookupError:
pass pass
self.process = None self.process = None
if self._config_path and self._config_path.exists():
self._config_path.unlink(missing_ok=True)
self._config_path = None
self._started = False self._started = False
# Tear down any owned upstream workers. # Tear down any owned upstream workers.
for w in self._owned_workers: for w in self._owned_workers:
@@ -292,77 +286,53 @@ class Gateway:
# ----- internals ------------------------------------------------------ # ----- internals ------------------------------------------------------
def _build_config( def _build_args(
self, self,
*, *,
model_id: str, model_id: str,
tokenizer_path: str, tokenizer_path: str,
urls: list[str], urls: list[str],
policy: str, policy: str,
extra_models: list[dict], ) -> list[str]:
) -> str:
resolved_tokenizer = _resolve_tokenizer_path(tokenizer_path) resolved_tokenizer = _resolve_tokenizer_path(tokenizer_path)
extra_model_toml = "" args = [
for em in extra_models: "--host",
extra_model_toml += ( self.host,
f'\n[[models]]\nid = "{em["id"]}"\n' "--port",
f'tokenizer_path = "{_resolve_tokenizer_path(em["tokenizer_path"])}"\n' str(self.port),
f'policy = "{em.get("policy", policy)}"\n' "--model-id",
) model_id,
"--tokenizer-path",
# Optional tunables — only emit the [proxy] and [active_load] resolved_tokenizer,
# sections if a test has overridden them, so production defaults "--policy",
# apply otherwise. policy,
proxy_section = "" ]
# Optional tunables — only pass them if a test overrode them, so
# the router's production defaults apply otherwise.
if self.proxy_request_timeout_secs is not None: if self.proxy_request_timeout_secs is not None:
proxy_section = ( args += ["--request-timeout-secs", str(self.proxy_request_timeout_secs)]
f"\n[proxy]\nrequest_timeout_secs = {self.proxy_request_timeout_secs}\n"
)
active_load_section = ""
if self.stale_request_timeout_secs is not None: if self.stale_request_timeout_secs is not None:
active_load_section = ( args += [
f"\n[active_load]\nstale_request_timeout_secs = " "--stale-request-timeout-secs",
f"{self.stale_request_timeout_secs}\n" str(self.stale_request_timeout_secs),
) ]
# `--worker-urls` is multi-valued; keep it last so clap doesn't
# absorb a following flag as a URL.
args += ["--worker-urls", *urls]
return args
urls_toml = ", ".join(f'"{u}"' for u in urls) def _launch(self, args: list[str], *, timeout: float) -> None:
return f"""\
[server]
host = "{self.host}"
port = {self.port}
[[models]]
id = "{model_id}"
tokenizer_path = "{resolved_tokenizer}"
policy = "{policy}"
{extra_model_toml}
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = [{urls_toml}]
{proxy_section}{active_load_section}"""
def _launch(self, config_text: str, *, timeout: float) -> None:
if not self.binary.exists(): if not self.binary.exists():
raise RuntimeError( raise RuntimeError(
f"sgl-router binary not found at {self.binary}. " f"sgl-router binary not found at {self.binary}. "
"Build it first: `cd experimental/sgl-router && cargo build --release` " "Build it first: `cd experimental/sgl-router && cargo build --release` "
"or set SGL_ROUTER_BINARY to the binary path." "or set SGL_ROUTER_BINARY to the binary path."
) )
# Write the main config. logger.info("sgl-router args: %s", args)
fd, path = tempfile.mkstemp(suffix=".toml", prefix="sgl-router-")
os.close(fd)
self._config_path = Path(path)
self._config_path.write_text(config_text, encoding="utf-8")
logger.info("sgl-router config: %s", self._config_path)
logger.debug("sgl-router config text:\n%s", config_text)
self.process = subprocess.Popen( self.process = subprocess.Popen(
[str(self.binary), "--config", str(self._config_path)], [str(self.binary), *args],
stdout=subprocess.PIPE, stdout=subprocess.PIPE,
stderr=subprocess.STDOUT, stderr=subprocess.STDOUT,
start_new_session=True, start_new_session=True,
@@ -380,16 +350,19 @@ urls = [{urls_toml}]
last_exc: Exception | None = None last_exc: Exception | None = None
while time.time() < deadline: while time.time() < deadline:
if self.process is not None and self.process.poll() is not None: if self.process is not None and self.process.poll() is not None:
# Process exited early — surface stdout/stderr. # Process exited early — surface stdout/stderr. This is the
out = b"" # primary startup-failure diagnostic, so if the read itself
# fails, report that instead of blanking the output.
try: try:
out = b""
if self.process.stdout is not None: if self.process.stdout is not None:
out = self.process.stdout.read() or b"" out = self.process.stdout.read() or b""
except Exception: # noqa: BLE001 output = out.decode(errors="replace")
pass except Exception as read_exc: # noqa: BLE001
output = f"<failed to read router stdout: {read_exc}>"
raise RuntimeError( raise RuntimeError(
f"sgl-router exited during startup with code " f"sgl-router exited during startup with code "
f"{self.process.returncode}. output:\n{out.decode(errors='replace')}", f"{self.process.returncode}. output:\n{output}",
) )
try: try:
resp = httpx.get(f"{self.base_url}/readyz", timeout=2.0) resp = httpx.get(f"{self.base_url}/readyz", timeout=2.0)
@@ -20,9 +20,23 @@ spec:
- name: router - name: router
image: sgl-router:e2e image: sgl-router:e2e
imagePullPolicy: Never imagePullPolicy: Never
# Configured entirely via CLI flags. No --service-discovery-namespace
# means a cluster-wide EndpointSlice watch (all namespaces); the
# `cross-ns-test=true` selector term scopes it to this test's workers.
args: args:
- "--config" - "--host"
- "/etc/config/router-cluster.toml" - "0.0.0.0"
- "--port"
- "8091"
- "--model-id"
- "tiny"
- "--tokenizer-path"
- "/etc/tokenizer/tiny.json"
- "--policy"
- "round_robin"
- "--service-discovery"
- "--selector"
- "app=sglang,cross-ns-test=true"
ports: ports:
- containerPort: 8091 - containerPort: 8091
name: http name: http
@@ -38,13 +52,6 @@ spec:
port: 8091 port: 8091
initialDelaySeconds: 5 initialDelaySeconds: 5
periodSeconds: 10 periodSeconds: 10
volumeMounts:
- name: config
mountPath: /etc/config
volumes:
- name: config
configMap:
name: sgl-router-cluster-config
--- ---
apiVersion: v1 apiVersion: v1
kind: Service kind: Service
@@ -18,9 +18,32 @@ spec:
- name: router - name: router
image: sgl-router:e2e image: sgl-router:e2e
imagePullPolicy: Never imagePullPolicy: Never
# Configured entirely via CLI flags. K8s EndpointSlice discovery
# watches `app=sglang` pods in this namespace. The aggressive
# circuit breaker (threshold 1, 5s cool-down) lets a terminating
# pod's connection-refused immediately drop it from the candidate
# set — the reconciliation tests scale workers rapidly and depend
# on fast eviction to absorb the churn.
args: args:
- "--config" - "--host"
- "/etc/config/router.toml" - "0.0.0.0"
- "--port"
- "8090"
- "--model-id"
- "tiny"
- "--tokenizer-path"
- "/etc/tokenizer/tiny.json"
- "--policy"
- "round_robin"
- "--cb-threshold"
- "1"
- "--cb-cool-down-secs"
- "5"
- "--service-discovery"
- "--service-discovery-namespace"
- "sgl-router-test"
- "--selector"
- "app=sglang"
ports: ports:
- containerPort: 8090 - containerPort: 8090
name: http name: http
@@ -36,13 +59,6 @@ spec:
port: 8090 port: 8090
initialDelaySeconds: 5 initialDelaySeconds: 5
periodSeconds: 10 periodSeconds: 10
volumeMounts:
- name: config
mountPath: /etc/config
volumes:
- name: config
configMap:
name: sgl-router-config
--- ---
apiVersion: v1 apiVersion: v1
kind: Service kind: Service
@@ -141,38 +141,9 @@ log "Waiting for fake-worker rollout..."
kubectl --context "${CONTEXT}" -n "${NAMESPACE}" rollout status deployment/fake-worker --timeout=120s kubectl --context "${CONTEXT}" -n "${NAMESPACE}" rollout status deployment/fake-worker --timeout=120s
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Step 6: Create sgl-router ConfigMap with k8s discovery pointing at the # Step 6: Deploy sgl-router. It is configured entirely via CLI flags in
# namespace where fake-worker pods live. # router.yaml — k8s EndpointSlice discovery watches `app=sglang`
# --------------------------------------------------------------------------- # pods in the sgl-router-test namespace (where fake-worker lives).
log "Creating sgl-router-config ConfigMap..."
ROUTER_CONFIG="[server]
host = \"0.0.0.0\"
port = 8090
[[models]]
id = \"tiny\"
tokenizer_path = \"/etc/tokenizer/tiny.json\"
policy = \"round_robin\"
# Aggressive breaker so a terminating pod's connection-refused
# immediately excludes it from the next request's candidate set —
# the reconciliation tests scale workers rapidly and depend on
# fast worker eviction to absorb the churn.
circuit_breaker = { threshold = 1, cool_down_secs = 5 }
[discovery]
backend = \"k8s\"
[discovery.k8s]
namespace = \"${NAMESPACE}\"
label_selector = \"app=sglang\""
kubectl --context "${CONTEXT}" -n "${NAMESPACE}" create configmap sgl-router-config \
--from-literal=router.toml="${ROUTER_CONFIG}" \
--dry-run=client -o yaml \
| kubectl --context "${CONTEXT}" apply -f -
# ---------------------------------------------------------------------------
# Step 7: Deploy sgl-router
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
log "Deploying sgl-router..." log "Deploying sgl-router..."
kubectl --context "${CONTEXT}" apply -f "${MANIFESTS_DIR}/router.yaml" kubectl --context "${CONTEXT}" apply -f "${MANIFESTS_DIR}/router.yaml"
@@ -149,49 +149,9 @@ def cluster_scoped_router(k8s_cluster):
_ensure_namespace(EXTRA_NAMESPACE) _ensure_namespace(EXTRA_NAMESPACE)
_ensure_service_in_ns(EXTRA_NAMESPACE) _ensure_service_in_ns(EXTRA_NAMESPACE)
# ConfigMap for the cluster-scoped router: empty namespace = watch all # The cluster-scoped router is configured via CLI flags in
cluster_config = """[server] # router-cluster-scoped.yaml: no --service-discovery-namespace (watch
host = "0.0.0.0" # all namespaces) and --selector app=sglang,cross-ns-test=true.
port = 8091
[[models]]
id = "tiny"
tokenizer_path = "/etc/tokenizer/tiny.json"
policy = "round_robin"
[discovery]
backend = "k8s"
[discovery.k8s]
namespace = ""
label_selector = "app=sglang,cross-ns-test=true"
"""
_kubectl(
"create",
"configmap",
"sgl-router-cluster-config",
f"--from-literal=router-cluster.toml={cluster_config}",
"-n",
NAMESPACE,
"--dry-run=client",
"-o",
"yaml",
check=True,
)
# pipe through apply
proc = _kubectl(
"create",
"configmap",
"sgl-router-cluster-config",
f"--from-literal=router-cluster.toml={cluster_config}",
"-n",
NAMESPACE,
"--dry-run=client",
"-o",
"yaml",
)
_apply_from_stdin(proc.stdout)
_kubectl("apply", "-f", str(router_manifest)) _kubectl("apply", "-f", str(router_manifest))
# The cluster-scoped router's /readyz blocks on registry-not-empty, so # The cluster-scoped router's /readyz blocks on registry-not-empty, so
@@ -2,8 +2,8 @@
// SPDX-License-Identifier: Apache-2.0 // SPDX-License-Identifier: Apache-2.0
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry; use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry;
@@ -29,18 +29,16 @@ fn config_for(_worker_url: &str) -> Config {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }
@@ -619,7 +617,7 @@ async fn no_healthy_workers_returns_503() {
); );
} }
/// A worker is registered for a model that is NOT in `cfg.models` (so the /// A worker is registered for a model that is NOT the configured `cfg.model` (so the
/// policy registry has no entry for it). The handler returns 404 /// policy registry has no entry for it). The handler returns 404
/// `model_not_found` rather than 500 — clients can recover by sending a /// `model_not_found` rather than 500 — clients can recover by sending a
/// different model name; an internal_error would mask the misconfiguration. /// different model name; an internal_error would mask the misconfiguration.
@@ -31,7 +31,7 @@ async fn failover_when_one_worker_dies() {
port: 0, port: 0,
}, },
observability: Default::default(), observability: Default::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
@@ -40,12 +40,10 @@ async fn failover_when_one_worker_dies() {
cool_down_secs: 30, cool_down_secs: 30,
}), }),
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec![w1.url.clone(), w2.url.clone(), w3.url.clone()], urls: vec![w1.url.clone(), w2.url.clone(), w3.url.clone()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
}; };
@@ -16,8 +16,8 @@
use bytes::Bytes; use bytes::Bytes;
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults; use sgl_router::policies::factory::build_registry_with_defaults;
@@ -40,18 +40,16 @@ fn build_ctx_with_worker(worker_url: &str) -> Arc<AppContext> {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
}; };
@@ -4,8 +4,8 @@
use axum::body::Body; use axum::body::Body;
use axum::http::Request; use axum::http::Request;
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry; use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry;
@@ -27,18 +27,16 @@ async fn forwards_whitelisted_headers_strips_others() {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
}; };
@@ -21,8 +21,8 @@ use axum::http::{Request, StatusCode};
use bytes::Bytes; use bytes::Bytes;
use serde_json::{json, Value}; use serde_json::{json, Value};
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults; use sgl_router::policies::factory::build_registry_with_defaults;
@@ -42,18 +42,16 @@ fn config() -> Config {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }
@@ -20,8 +20,8 @@ use axum::body::Body;
use axum::http::{Request, StatusCode}; use axum::http::{Request, StatusCode};
use http_body_util::BodyExt; use http_body_util::BodyExt;
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults; use sgl_router::policies::factory::build_registry_with_defaults;
@@ -41,18 +41,16 @@ fn config() -> Config {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }
@@ -13,8 +13,8 @@ use axum::body::Body;
use axum::http::{Request, StatusCode}; use axum::http::{Request, StatusCode};
use http_body_util::BodyExt; use http_body_util::BodyExt;
use sgl_router::config::{ use sgl_router::config::{
ActiveLoadConfig, Config, DiscoveryBackend, DiscoveryConfig, ModelConfig, ObservabilityConfig, ActiveLoadConfig, Config, DiscoveryBackend, ModelConfig, ObservabilityConfig, PolicyKind,
PolicyKind, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig, ProxyConfig, ServerConfig, StaticUrlsDiscoveryConfig,
}; };
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec}; use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry; use sgl_router::policies::factory::build_registry_with_defaults as build_policy_registry;
@@ -34,18 +34,16 @@ fn config(_worker_url: &str) -> Config {
port: 0, port: 0,
}, },
observability: ObservabilityConfig::default(), observability: ObservabilityConfig::default(),
models: vec![ModelConfig { model: ModelConfig {
id: "tiny".into(), id: "tiny".into(),
tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(), tokenizer_path: "tests/fixtures/tiny_tokenizer.json".into(),
policy: PolicyKind::RoundRobin, policy: PolicyKind::RoundRobin,
circuit_breaker: None, circuit_breaker: None,
cache_aware: None, cache_aware: None,
}], },
discovery: DiscoveryConfig { discovery: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
backend: DiscoveryBackend::StaticUrls(StaticUrlsDiscoveryConfig {
urls: vec!["http://placeholder:0".into()], urls: vec!["http://placeholder:0".into()],
}), }),
},
proxy: ProxyConfig::default(), proxy: ProxyConfig::default(),
active_load: ActiveLoadConfig::default(), active_load: ActiveLoadConfig::default(),
} }