Rust server: align launcher and request validation behavior (#37327)

This commit is contained in:
Lianmin Zheng
2026-09-01 23:42:36 -07:00
committed by GitHub
parent 83bd2c473f
commit f8f04bafa8
13 changed files with 317 additions and 41 deletions
@@ -196,7 +196,7 @@ async fn generate(
State(state): State<Arc<AppState>>,
body: Result<Json<GenerateBody>, JsonRejection>,
) -> Response {
let body = match body {
let mut body = match body {
Ok(Json(body)) => body,
// A body that fails to parse has no readable `stream` flag, so this one
// can only answer unary — as Python's does (FastAPI rejects before its
@@ -206,6 +206,15 @@ async fn generate(
}
};
let stream = body.stream;
if let Some(preferred) = &state.server_args.preferred_sampling_params
&& let Err(error) = body.apply_preferred_sampling(&preferred.0)
{
return native_error(
StatusCode::INTERNAL_SERVER_ERROR,
&error.to_string(),
stream,
);
}
// Fan `text`/`input_ids`/`sampling_params` (scalar or list) into per-request
// payloads. `is_batch` = list form → the response is a JSON array.
let (mut payloads, is_batch) = match body.into_requests() {
+6 -2
View File
@@ -126,8 +126,8 @@ pub struct ServerArgs {
pub disaggregation_mode: DisaggregationMode,
/// The resolved Python `ModelConfig`, attached at handoff time.
pub model_config: ModelConfig,
/// Default sampling params advertised by `/get_model_info`, verbatim from
/// `server_args.preferred_sampling_params` (a JSON object or null).
/// Launch-time sampling defaults merged beneath per-request values and
/// advertised by `/get_model_info`.
pub preferred_sampling_params: Option<PreferredSamplingParams>,
/// Over-long inputs are truncated to fit the context instead of 400ing, and
/// `max_new_tokens` is clamped rather than rejected (Python
@@ -528,6 +528,10 @@ impl ServerArgs {
if self.served_model_name.is_empty() {
return Err("empty 'served_model_name' in server_args".into());
}
if let Some(preferred) = &self.preferred_sampling_params {
super::sampling::SamplingParamsInput::from_preferred(&preferred.0)
.map_err(|e| format!("invalid preferred_sampling_params: {e}"))?;
}
Ok(())
}
+12
View File
@@ -109,6 +109,18 @@ pub struct GenerateBody {
}
impl GenerateBody {
/// Merge operator-provided sampling defaults beneath request values,
/// matching Python TokenizerManager's preferred/request precedence.
pub fn apply_preferred_sampling(&mut self, preferred: &serde_json::Value) -> Result<(), Error> {
match &mut self.sampling_params {
Some(params) => params.apply_preferred(preferred),
None => SamplingParamsInput::from_preferred(preferred).map(|params| {
self.sampling_params = Some(params);
}),
}
.map_err(|e| Error::Validation(format!("invalid preferred_sampling_params: {e}")))
}
/// Validate, normalize and fan the body into one [`GenerateRequest`] per
/// prompt + `is_batch` (list form — a 1-element list is still a batch → JSON
/// array response). The Rust counterpart of Python
+99 -3
View File
@@ -3,7 +3,7 @@
//! `__post_init__` → `normalize` → `verify` pipeline (run in that order, as
//! `TokenizerManager._create_tokenized_object` does).
use std::collections::BTreeMap;
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use serde::de::value::{MapAccessDeserializer, SeqAccessDeserializer};
@@ -233,6 +233,11 @@ pub struct SamplingParams {
/// Set by `normalize`; tells the scheduler its own pass can early-return.
#[serde(skip_deserializing)]
pub is_normalized: bool,
/// API fields present in the request object. Serde defaults erase this
/// distinction, but preferred sampling parameters must not overwrite an
/// explicit request value, including an explicit default or null.
#[serde(skip)]
pub(crate) explicit_fields: BTreeSet<String>,
}
/// The `/generate` body's `sampling_params`: one object (broadcast to every
@@ -263,12 +268,21 @@ impl<'de> Deserialize<'de> for SamplingParamsInput {
}
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Self::Value, A::Error> {
SamplingParams::deserialize(MapAccessDeserializer::new(map))
let value = serde_json::Value::deserialize(MapAccessDeserializer::new(map))?;
sampling_params_from_value(value)
.map(|p| SamplingParamsInput::One(Box::new(p)))
.map_err(serde::de::Error::custom)
}
fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<Self::Value, A::Error> {
Vec::deserialize(SeqAccessDeserializer::new(seq)).map(SamplingParamsInput::Many)
let values =
Vec::<serde_json::Value>::deserialize(SeqAccessDeserializer::new(seq))?;
values
.into_iter()
.map(sampling_params_from_value)
.collect::<Result<Vec<_>, _>>()
.map(SamplingParamsInput::Many)
.map_err(serde::de::Error::custom)
}
}
@@ -276,6 +290,56 @@ impl<'de> Deserialize<'de> for SamplingParamsInput {
}
}
fn sampling_params_from_value(value: serde_json::Value) -> Result<SamplingParams, String> {
let explicit_fields = value
.as_object()
.ok_or_else(|| "sampling_params must be an object".to_string())?
.keys()
.cloned()
.collect();
let mut params: SamplingParams = serde_json::from_value(value).map_err(|e| e.to_string())?;
params.explicit_fields = explicit_fields;
Ok(params)
}
impl SamplingParamsInput {
/// Merge launch-time preferred params beneath request params. A request key
/// wins even when it explicitly carries the type's default or null.
pub fn apply_preferred(&mut self, preferred: &serde_json::Value) -> Result<(), String> {
match self {
Self::One(params) => apply_preferred_to_one(params, preferred),
Self::Many(params) => params
.iter_mut()
.try_for_each(|params| apply_preferred_to_one(params, preferred)),
}
}
pub fn from_preferred(preferred: &serde_json::Value) -> Result<Self, String> {
sampling_params_from_value(preferred.clone()).map(|params| Self::One(Box::new(params)))
}
}
fn apply_preferred_to_one(
params: &mut SamplingParams,
preferred: &serde_json::Value,
) -> Result<(), String> {
let mut merged = preferred
.as_object()
.ok_or_else(|| "preferred_sampling_params must be a JSON object".to_string())?
.clone();
let request_value = serde_json::to_value(&*params).map_err(|e| e.to_string())?;
let request = request_value
.as_object()
.ok_or_else(|| "SamplingParams did not serialize as an object".to_string())?;
for field in &params.explicit_fields {
if let Some(value) = request.get(field) {
merged.insert(field.clone(), value.clone());
}
}
*params = sampling_params_from_value(serde_json::Value::Object(merged))?;
Ok(())
}
impl Default for SamplingParams {
fn default() -> Self {
// Each field reads the same `default()` the serde attribute above names,
@@ -314,6 +378,7 @@ impl Default for SamplingParams {
stop_str_max_len: 0,
stop_regex_max_len: 0,
is_normalized: false,
explicit_fields: BTreeSet::new(),
}
}
}
@@ -1085,4 +1150,35 @@ mod tests {
let err = norm_err(&json).to_string();
assert!(err.contains("at most"), "{err}");
}
#[test]
fn preferred_params_fill_only_omitted_request_fields() {
let preferred = serde_json::json!({
"temperature": 0.25,
"top_p": 0.75,
"max_new_tokens": 4096
});
let mut input: SamplingParamsInput =
serde_json::from_str(r#"{"temperature": 1.0, "top_p": null}"#).unwrap();
input.apply_preferred(&preferred).unwrap();
let SamplingParamsInput::One(params) = input else {
panic!("expected scalar params")
};
assert_eq!(params.temperature, 1.0, "explicit default wins");
assert_eq!(params.top_p, 1.0, "explicit null keeps the type default");
assert_eq!(params.max_new_tokens, Some(4096), "omitted uses preferred");
}
#[test]
fn preferred_params_apply_to_every_batched_object() {
let preferred = serde_json::json!({"temperature": 0.25, "top_p": 0.75});
let mut input: SamplingParamsInput =
serde_json::from_str(r#"[{"temperature": 0.5}, {"top_p": 0.9}]"#).unwrap();
input.apply_preferred(&preferred).unwrap();
let SamplingParamsInput::Many(params) = input else {
panic!("expected batched params")
};
assert_eq!((params[0].temperature, params[0].top_p), (0.5, 0.75));
assert_eq!((params[1].temperature, params[1].top_p), (0.25, 0.9));
}
}
@@ -12,7 +12,9 @@ use crate::message::response::ResponseItem;
use crate::runtime::Runnable;
use crate::tokenizer_manager::channel::ToSchedulerTx;
pub use crate::tokenizer_manager::to_scheduler_types::{Limits, Mm};
use crate::tokenizer_manager::to_scheduler_validation::{check_total_tokens, validate};
use crate::tokenizer_manager::to_scheduler_validation::{
check_total_tokens, validate, validate_input_ids,
};
use crate::tokenizer_manager::wiring::{AbortSource, Senders, TmEvent};
use crate::utils::{
error::Error,
@@ -274,7 +276,8 @@ impl Intake {
// a text request has no ids yet.
RequestState::PreSendValidating => {
if let RequestKind::Generate(g) = &mut req.kind
&& let Err(e) = check_total_tokens(g, &self.limits)
&& let Err(e) = validate_input_ids(g, self.limits.vocab_size)
.and_then(|()| check_total_tokens(g, &self.limits))
{
let _ = req.state.apply(Event::Error(e)); // → Failed
continue;
@@ -636,6 +636,27 @@ fn negative_and_logprob_token_ids_rejected() {
}
}
#[test]
fn multimodal_sentinel_is_validated_after_expansion() {
let mut req = generate_req(24, SamplingParams::default());
let RequestKind::Generate(g) = &mut req.kind else {
unreachable!()
};
g.input_ids = Some(vec![1, -103, 2]);
g.mm = Some(Box::new(crate::message::request::MmData {
audio_data: Some(rmpv::Value::from("data:audio/wav;base64,xxxx")),
..Default::default()
}));
assert!(validate(&mut req, &test_limits()).is_ok());
let RequestKind::Generate(g) = &mut req.kind else {
unreachable!()
};
assert!(validate_input_ids(g, test_limits().vocab_size).is_err());
g.input_ids = Some(vec![1, 103, 2]);
assert!(validate_input_ids(g, test_limits().vocab_size).is_ok());
}
/// A valid request is registered and handed onward — never deregistered.
#[test]
fn admitted_request_keeps_registration() {
@@ -41,19 +41,12 @@ pub(super) fn validate(req: &mut Request, limits: &Limits) -> Result<(), Error>
));
}
// Client-supplied token ids must be in-vocabulary: an out-of-range id
// reaches the embedding lookup and kills the scheduler process, so 400
// here instead — mirroring the Python `TokenizerManager` validation.
// Multimodal processors may consume sentinel ids outside the vocabulary.
// Validate their resulting ids at PreSendValidating instead. Non-MM client
// ids can be rejected now.
if let RequestKind::Generate(g) = &req.kind {
if let Some(ids) = &g.input_ids {
for &id in ids {
if id < 0 || id as u64 >= vocab_size {
return Err(Error::Validation(format!(
"input_ids contains out-of-vocabulary token id {id}; \
valid range is [0, {vocab_size})"
)));
}
}
if !g.has_multimodal() {
validate_input_ids(g, vocab_size)?;
}
if let Some(ids) = &g.token_ids_logprob {
for &id in ids {
@@ -95,6 +88,22 @@ pub(super) fn validate(req: &mut Request, limits: &Limits) -> Result<(), Error>
Ok(())
}
/// Guard the ids that will reach the embedding lookup. Multimodal requests run
/// this after placeholder expansion; all other requests also run it at intake.
pub(super) fn validate_input_ids(g: &GenerateRequest, vocab_size: u64) -> Result<(), Error> {
if let Some(ids) = &g.input_ids {
for &id in ids {
if id < 0 || id as u64 >= vocab_size {
return Err(Error::Validation(format!(
"input_ids contains out-of-vocabulary token id {id}; \
valid range is [0, {vocab_size})"
)));
}
}
}
Ok(())
}
/// The context-window checks that need the tokenized length, mirroring Python
/// `TokenizerManager._validate_one_request`: the input alone must fit, and then
/// input + `max_new_tokens` must fit. Without them the scheduler silently clamps