Rust server: align launcher and request validation behavior (#37327)
This commit is contained in:
@@ -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() {
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ¶ms.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
|
||||
|
||||
Reference in New Issue
Block a user