[Router] Fleet-wide sampling contract 3/3: splice injection without re-serializing (#39002)

Co-authored-by: Kangyan Zhou <kangyan.zhou@radixark.ai>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-09-16 09:59:48 -07:00
committed by GitHub
co-authored by Kangyan Zhou Claude Opus 5
parent 3e03879f68
commit b02e16a895
4 changed files with 593 additions and 8 deletions
@@ -24,6 +24,7 @@ mod pd_pool_isolation;
mod pd_protocol_binding;
mod radix_tree_routing;
mod roundrobin_input_ids;
mod sampling_overrides;
mod shared_prefill_admission;
mod sticky_input_ids;
mod sticky_routing;
@@ -0,0 +1,314 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
//! `--override-sampling-params` end to end: from the CLI flag an operator
//! writes in a manifest to the JSON body the engine actually receives.
//!
//! The unit tests in `config::sampling` and `server::routes::chat` cover
//! parsing and the per-parameter decision; these drive the whole path, because
//! the failure this flag exists to prevent (an engine serving sampling
//! parameters the operator did not declare) is only observable on the wire.
use axum::body::Body;
use axum::http::{Request, StatusCode};
use serde_json::{json, Value};
use sgl_router::config::{Cli, Config};
use sgl_router::discovery::{ModelId, WorkerId, WorkerMode, WorkerSpec};
use sgl_router::policies::factory::build_registry_with_defaults;
use sgl_router::proxy::Proxy;
use sgl_router::server::app::build_router;
use sgl_router::server::app_context::AppContext;
use sgl_router::tokenizer::TokenizerRegistry;
use sgl_router::workers::WorkerRegistry;
use std::sync::Arc;
use std::time::Duration;
use tower::ServiceExt;
use crate::common::mock_worker::MockWorker;
const MODEL: &str = "tiny";
const OVERRIDES: &str = r#"{"temperature": 1, "top_p": 0.95, "top_k": 1000,
"frequency_penalty": 0, "presence_penalty": 0, "n": 1}"#;
/// Build the config the way a deployment does — through `Cli`, so what these
/// tests pin is the flag spelling in a manifest, not a hand-built struct that
/// could drift from what the parser produces.
fn config(flags: &[&str]) -> Config {
let mut argv = vec![
"sgl-router",
"--model-id",
MODEL,
"--tokenizer-path",
"tests/fixtures/tiny_tokenizer.json",
"--worker-urls",
"http://placeholder:0",
];
argv.extend_from_slice(flags);
<Cli as clap::Parser>::parse_from(argv)
.into_config()
.expect("flags must parse")
}
fn build_ctx(url: String, flags: &[&str]) -> Arc<AppContext> {
let cfg = config(flags);
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let registry = Arc::new(WorkerRegistry::default());
let _ = registry.add(WorkerSpec {
id: WorkerId(url.clone()),
url,
mode: WorkerMode::Plain,
model_ids: vec![ModelId(MODEL.into())],
bootstrap_port: None,
});
let policies = Arc::new(build_registry_with_defaults(&cfg).unwrap());
let proxy = Arc::new(Proxy::new(Duration::from_secs(5)).unwrap());
Arc::new(AppContext::new(cfg, tokenizers, proxy, registry, policies))
}
async fn send(ctx: Arc<AppContext>, body: Value) -> StatusCode {
send_raw(ctx, serde_json::to_vec(&body).unwrap()).await.0
}
/// Send a body verbatim, so a test can express what `serde_json::Value`
/// cannot — a repeated key. Returns the status and the router's own error
/// code header.
async fn send_raw(ctx: Arc<AppContext>, body: Vec<u8>) -> (StatusCode, Option<String>) {
let req = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(body))
.unwrap();
let resp = build_router(ctx).oneshot(req).await.unwrap();
let code = resp
.headers()
.get("x-router-error-code")
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
(resp.status(), code)
}
fn captured(mock: &MockWorker) -> Option<Value> {
let b = mock.captured.lock().unwrap().last_body.clone()?;
Some(serde_json::from_slice(&b).expect("captured body is valid JSON"))
}
/// The values an operator configures replace the engine's own defaults: a
/// request that names no sampling parameter reaches the engine carrying every
/// configured one.
#[tokio::test]
async fn configured_values_reach_the_engine_when_the_request_omits_them() {
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &["--override-sampling-params", OVERRIDES]);
let status = send(
ctx,
json!({"model": MODEL, "messages": [{"role": "user", "content": "hi"}]}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("worker received a request");
assert_eq!(body.get("temperature"), Some(&json!(1)));
assert_eq!(body.get("top_p"), Some(&json!(0.95)));
assert_eq!(body.get("top_k"), Some(&json!(1000)));
assert_eq!(body.get("frequency_penalty"), Some(&json!(0)));
assert_eq!(body.get("presence_penalty"), Some(&json!(0)));
assert_eq!(body.get("n"), Some(&json!(1)));
}
/// Under the default `reject` mode a conflicting request is a 400 that never
/// reaches a worker — the contract costs no engine round-trip and no
/// admission slot.
#[tokio::test]
async fn reject_mode_400s_a_conflicting_request_without_touching_the_engine() {
// Covers the whole rejection contract: status, wire error code, and that
// the engine is never reached.
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &["--override-sampling-params", OVERRIDES]);
let body = json!({
"model": MODEL,
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.7,
});
let (status, code) = send_raw(ctx, serde_json::to_vec(&body).unwrap()).await;
assert_eq!(status, StatusCode::BAD_REQUEST);
// Its own error code, so an operator rolling `reject` across a fleet can
// alert on contract rejections without them being buried among the
// `bad_request`s from clients sending malformed JSON.
assert_eq!(code.as_deref(), Some("sampling_contract_violation"));
assert!(
captured(&mock).is_none(),
"a rejected request must not reach the engine"
);
}
/// A repeated sampling key must not become a router-side 400: it is legal JSON
/// that every engine reads last-wins, and the router forwarded it before these
/// fields were probed. The contract judges the value the engine will use.
#[tokio::test]
async fn duplicate_sampling_key_is_judged_on_its_last_value() {
// Last value agrees with the contract -> served.
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &["--override-sampling-params", OVERRIDES]);
let (status, _) = send_raw(
ctx,
br#"{"model":"tiny","messages":[{"role":"user","content":"hi"}],
"temperature":0.7,"temperature":1}"#
.to_vec(),
)
.await;
assert_eq!(status, StatusCode::OK);
assert_eq!(
captured(&mock).and_then(|b| b.get("temperature").cloned()),
Some(json!(1)),
"the engine must see the last value"
);
// Last value disagrees -> rejected, even though the first one matched.
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &["--override-sampling-params", OVERRIDES]);
let (status, _) = send_raw(
ctx,
br#"{"model":"tiny","messages":[{"role":"user","content":"hi"}],
"temperature":1,"temperature":0.7}"#
.to_vec(),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST);
}
/// `temperature`, `top_p`, `top_k`, `min_p` and `repetition_penalty` are the
/// parameters the engine resolves from the model's own `generation_config`, so
/// they are the ones a fleet-wide pin exists for. Drive them the whole way to
/// the engine.
#[tokio::test]
async fn engine_defaulted_parameters_reach_the_engine() {
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(
mock.url.clone(),
&[
"--override-sampling-params",
r#"{"min_p": 0.05, "repetition_penalty": 1.1}"#,
],
);
let status = send(
ctx,
json!({"model": MODEL, "messages": [{"role": "user", "content": "hi"}]}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("engine must receive a body");
assert_eq!(body.get("min_p"), Some(&json!(0.05)));
assert_eq!(body.get("repetition_penalty"), Some(&json!(1.1)));
}
/// The same request under `allow` is forwarded with the client's value intact,
/// and the parameters it did not name still get the configured ones.
#[tokio::test]
async fn allow_mode_forwards_the_client_value_and_fills_the_rest() {
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(
mock.url.clone(),
&[
"--override-sampling-params",
OVERRIDES,
"--sampling-param-conflict",
"allow",
],
);
let status = send(
ctx,
json!({
"model": MODEL,
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.7,
}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("worker received a request");
assert_eq!(body.get("temperature"), Some(&json!(0.7)));
assert_eq!(body.get("top_p"), Some(&json!(0.95)));
assert_eq!(body.get("top_k"), Some(&json!(1000)));
assert_eq!(body.get("n"), Some(&json!(1)));
}
/// A band accepts anything inside it, 400s outside, and injects nothing:
/// temperature tunable in [0, 1], everything else fixed.
#[tokio::test]
async fn a_temperature_band_admits_in_range_values_and_injects_nothing() {
let flags = [
"--override-sampling-params",
r#"{"temperature": {"min": 0, "max": 1}, "top_p": 0.95}"#,
];
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &flags);
let status = send(
ctx,
json!({
"model": MODEL,
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.6,
}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("worker received a request");
assert_eq!(body.get("temperature"), Some(&json!(0.6)));
assert_eq!(body.get("top_p"), Some(&json!(0.95)));
// Omitted: the band names no value, so the engine's own default applies.
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &flags);
let status = send(
ctx,
json!({"model": MODEL, "messages": [{"role": "user", "content": "hi"}]}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("worker received a request");
assert_eq!(body.get("temperature"), None);
// Outside the band: 400.
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &flags);
let status = send(
ctx,
json!({
"model": MODEL,
"messages": [{"role": "user", "content": "hi"}],
"temperature": 1.5,
}),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST);
assert!(captured(&mock).is_none());
}
/// With the flag unset the router touches nothing: the body the client sent
/// is the body the engine sees, including a sampling parameter the operator
/// could have fixed.
#[tokio::test]
async fn unset_flag_forwards_the_body_untouched() {
let mock = MockWorker::start(vec![]).await;
let ctx = build_ctx(mock.url.clone(), &[]);
let status = send(
ctx,
json!({
"model": MODEL,
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.7,
}),
)
.await;
assert_eq!(status, StatusCode::OK);
let body = captured(&mock).expect("worker received a request");
assert_eq!(body.get("temperature"), Some(&json!(0.7)));
assert_eq!(body.get("top_p"), None);
assert_eq!(body.get("top_k"), None);
assert_eq!(body.get("n"), None);
}