From 2ca2ca753a1bc7bd79707108f84b0479501b2e5e Mon Sep 17 00:00:00 2001 From: Rain Jiang <96632942+rainj-me@users.noreply.github.com> Date: Wed, 29 Jul 2026 10:54:56 -0700 Subject: [PATCH] sglang rust server request message (#32242) --- rust/sglang-server/src/ids.rs | 271 +++++-- rust/sglang-server/src/lib.rs | 1 + rust/sglang-server/src/message.rs | 18 + rust/sglang-server/src/message/io_struct.rs | 173 +++++ rust/sglang-server/src/message/request.rs | 820 ++++++++++++++++++++ rust/sglang-server/src/message/sampling.rs | 80 ++ rust/sglang-server/src/message/types.rs | 265 +++++++ 7 files changed, 1579 insertions(+), 49 deletions(-) create mode 100644 rust/sglang-server/src/message.rs create mode 100644 rust/sglang-server/src/message/io_struct.rs create mode 100644 rust/sglang-server/src/message/request.rs create mode 100644 rust/sglang-server/src/message/sampling.rs create mode 100644 rust/sglang-server/src/message/types.rs diff --git a/rust/sglang-server/src/ids.rs b/rust/sglang-server/src/ids.rs index 5f3e9f0cc..3b3ba9640 100644 --- a/rust/sglang-server/src/ids.rs +++ b/rust/sglang-server/src/ids.rs @@ -1,62 +1,175 @@ //! Lightweight identifiers used across pipeline stages. -#![allow(dead_code)] // TODO: remove when the consumer PR lands use std::{ - collections::hash_map::DefaultHasher, - fmt::{Debug, Display, Formatter, Result}, - hash::{Hash, Hasher}, + collections::hash_map::RandomState, + fmt, + hash::{BuildHasher, Hash}, + ops::Deref, + sync::OnceLock, }; use uuid::Uuid; -/// Process-local request key, derived by hashing the client-visible rid string -/// (`from_rid`). Cheap to copy, used as the routing key on the egress side -/// (detok shard selection) and to correlate scheduler output chunks back to the -/// originating connection. NOT the identity the client sees — that is the rid -/// string; this is its stable 64-bit digest. -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] -pub struct RidHash(pub u64); - -impl RidHash { - /// Derive the routing key from a rid string. `DefaultHasher::new()` uses - /// fixed keys, so every stage (ingress push, egress decode, abort) computes - /// the same id from the same rid — no shared map. - pub fn from_rid(rid: &str) -> Self { - let mut h = DefaultHasher::new(); - rid.hash(&mut h); - RidHash(h.finish()) - } - - /// Shard index for `n` detokenizer shards. Pure function of the id so the - /// ingress and egress sides agree without any shared map. - #[inline] - pub fn shard(self, n: usize) -> usize { - debug_assert!(n > 0); - (self.0 as usize) % n - } -} - -impl Display for RidHash { - fn fmt(&self, f: &mut Formatter<'_>) -> Result { - write!(f, "req-{}", self.0) - } -} - -/// Mint a fresh client-visible request id — uuid4 hex, matching the Python -/// server's `uuid.uuid4().hex`. -pub fn new_rid() -> String { - Uuid::new_v4().simple().to_string() -} - /// Health-probe rid prefix — MUST match the Python server's /// `sglang.srt.constants.HEALTH_CHECK_RID_PREFIX`, so scheduler logs / crash /// dumps and any prefix-gated logic recognize probes from either server. pub const HEALTH_CHECK_RID_PREFIX: &str = "HEALTH_CHECK"; -/// Mint a health-probe rid: `HEALTH_CHECK_`, the Python server's -/// `f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}"` format. -pub fn new_health_check_rid() -> String { - format!("{HEALTH_CHECK_RID_PREFIX}_{}", new_rid()) +/// Separates a client-supplied rid from the uniquifier appended to it (see +/// [`Rid::from_client`]). Deliberately a character this server never mints: +/// [`Rid::new`] is uuid hex and [`Rid::new_health_check`] adds only +/// `HEALTH_CHECK_`, so its presence at the fixed offset below is what lets +/// [`Rid::client_facing`] recognize a suffix without `Rid` carrying a flag — +/// which matters because `Rid` rides on every `ChunkEvent`. +const UNIQ_SEP: u8 = b'#'; +/// Hex digits of uniquifier: 8 of a per-process random base, 8 of a counter. +const UNIQ_DIGITS: usize = 16; +/// Total bytes appended. Fixed-width by construction — both halves are `u32` +/// formatted `{:08x}` — which is what makes stripping a slice, not a search. +const UNIQ_SUFFIX_LEN: usize = 1 + UNIQ_DIGITS; + +#[derive(Clone, Debug)] +pub struct Rid { + id: String, + /// Partition key, derived from `id`. Never part of identity — see the `Eq` / + /// `Hash` impls below. + hash: u64, +} + +// Identity is the ID, not the digest. Deriving these would fold `hash` into both, +// which is redundant while the seed is stable and silently wrong if it ever isn't. +impl PartialEq for Rid { + fn eq(&self, other: &Self) -> bool { + self.id == other.id + } +} +impl Eq for Rid {} +impl Hash for Rid { + fn hash(&self, state: &mut H) { + self.id.hash(state); + } +} + +impl Rid { + /// Borrow the underlying id — the form the msgpack wire and the HTTP body use. + pub fn as_str(&self) -> &str { + &self.id + } + + pub fn new() -> Self { + let id = Uuid::new_v4().simple().to_string(); + Rid::from(id) + } + + pub fn new_health_check() -> Self { + let id = format!("{HEALTH_CHECK_RID_PREFIX}_{}", Uuid::new_v4().simple()); + Rid::from(id) + } + + /// A CLIENT-SUPPLIED rid, made unique for internal use by appending a + /// uniquifier. + /// + /// Nothing stops two concurrent requests from arriving with the same rid, and + /// the rid is an identity downstream: detok `Register` is an insert-overwrite, + /// so the second would evict the first's sink, 500 that client mid-generation + /// and deliver its remaining chunks to the second's connection. Uniquifying + /// here makes the collision unrepresentable rather than something a duplicate + /// check has to catch — no in-flight registry, no admission/release ordering + /// to get wrong, and both clients get served instead of the second being 400'd. + /// + /// The client never sees this: [`client_facing`](Self::client_facing) strips it + /// back off for `meta_info.id`. Only the scheduler wire and its logs carry the + /// suffixed form. + /// + /// The counter alone would guarantee uniqueness within a process; the random + /// base covers several HTTP worker processes feeding one scheduler, where two + /// counters would otherwise both start at zero. `u32` (not `u64`) keeps the + /// `{:08x}` width exactly 8 — wrapping needs 2^32 live requests sharing one + /// client rid. + pub fn from_client(id: &str) -> Self { + use std::sync::atomic::{AtomicU32, Ordering}; + static BASE: OnceLock = OnceLock::new(); + static NEXT: AtomicU32 = AtomicU32::new(0); + let base = *BASE.get_or_init(|| Uuid::new_v4().as_u128() as u32); + let n = NEXT.fetch_add(1, Ordering::Relaxed); + Rid::from(format!( + "{id}{sep}{base:08x}{n:08x}", + sep = UNIQ_SEP as char + )) + } + + /// The rid as the client wrote it — what `meta_info.id` must echo, and what + /// the client-facing length limit applies to. Returns the whole id for a rid + /// this server minted, which carries no suffix to strip. + pub fn client_facing(&self) -> &str { + let b = self.id.as_bytes(); + let Some(cut) = b.len().checked_sub(UNIQ_SUFFIX_LEN) else { + return &self.id; + }; + // `UNIQ_SEP` is ASCII, so a match here is also a char boundary — slicing + // at `cut` cannot split a multi-byte character in a client's rid. + if b[cut] == UNIQ_SEP && b[cut + 1..].iter().all(u8::is_ascii_hexdigit) { + &self.id[..cut] + } else { + &self.id + } + } + + /// Shard index for `n` detokenizer shards. Pure function of the id so the + /// ingress and egress sides agree without any shared map. + #[inline] + pub fn shard(&self, n: usize) -> usize { + debug_assert!(n > 0); + (self.hash as usize) % n + } +} + +impl From for Rid { + fn from(id: String) -> Self { + // ONE seed per process, not one per conversion. Ingress and egress each + // build a `Rid` from the same string and must agree on the shard without a + // shared map — a fresh `RandomState` here would hash the same rid two + // different ways, so chunks would arrive at a shard that never registered + // the request and be dropped. + // + // The seed is random rather than fixed because rids are client-supplied: + // with public keys, colliding rids are an offline ~2^32 search. Collisions + // are only a shard co-location now (identity is the string), but a keyed + // hash also stops an attacker from stacking every request onto one shard. + static SEED: OnceLock = OnceLock::new(); + let hash = SEED.get_or_init(RandomState::new).hash_one(&id); + Rid { id, hash } + } +} + +impl From<&str> for Rid { + fn from(id: &str) -> Self { + Rid::from(id.to_string()) + } +} + +impl Deref for Rid { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.id + } +} + +impl Default for Rid { + fn default() -> Self { + Rid::new() + } +} + +impl fmt::Display for Rid { + /// The BARE rid, with no decoration. It is formatted into client-facing error + /// messages and into wire values (`AbortReq`), so a prefix here would surface + /// as a corrupted id rather than a nicety. `Debug` still shows `Rid("…")` if a + /// log wants the type visible. + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.id) + } } #[cfg(test)] @@ -69,7 +182,7 @@ mod tests { /// hyphenated) would silently break the parity. #[test] fn rid_matches_python_uuid4_hex_format() { - let rid = new_rid(); + let rid = Rid::new(); assert_eq!(rid.len(), 32); assert!( rid.chars() @@ -78,13 +191,73 @@ mod tests { ); } + /// The round-trip that makes the scheme invisible: whatever the client sent + /// comes back out of `client_facing`, byte for byte, however odd it is. + /// `meta_info.id` is the only thing the client can correlate a response by, so + /// leaking the uniquifier — or over-stripping a rid that happens to look like + /// one — is a client-visible bug. + #[test] + fn client_facing_round_trips_whatever_the_client_sent() { + for given in [ + "r", + "", + "x_0", + "0123456789abcdef0123456789abcdef", + // Already shaped like a suffix: stripping must remove OURS, not theirs. + "abc#0123456789abcdef", + "#0123456789abcdef", + // Non-ASCII, so a byte-offset slice could split a character. + "réq-π-🎉", + &"x".repeat(128), + ] { + let rid = Rid::from_client(given); + assert_eq!( + rid.client_facing(), + given, + "round-trip failed for {given:?}" + ); + assert_ne!(rid.as_str(), given, "the internal rid must be uniquified"); + } + } + + /// A rid this server minted carries no suffix, so `client_facing` must return + /// it whole. Nothing strips these today, but `Rid::new` is uuid hex and a + /// 17-byte tail of it is all hex — only the missing separator saves it. + #[test] + fn client_facing_leaves_minted_rids_alone() { + for rid in [Rid::new(), Rid::new_health_check(), Rid::default()] { + assert_eq!(rid.client_facing(), rid.as_str()); + } + } + + /// Uniqueness is the entire point, and it must hold for the same input — that + /// IS the collision case. Checked across threads because the counter is shared + /// by every api thread. + #[test] + fn from_client_is_unique_even_for_one_repeated_rid() { + let handles: Vec<_> = (0..4) + .map(|_| { + std::thread::spawn(|| { + (0..250) + .map(|_| Rid::from_client("same").as_str().to_string()) + .collect::>() + }) + }) + .collect(); + let all: std::collections::HashSet = handles + .into_iter() + .flat_map(|h| h.join().unwrap()) + .collect(); + assert_eq!(all.len(), 1000, "every uniquified rid must be distinct"); + } + /// Cross-language literal guard: the prefix is dictated by Python's /// `constants.HEALTH_CHECK_RID_PREFIX` ("HEALTH_CHECK"); drifting silently /// would break prefix-gated handling (e.g. the disagg encode server). #[test] fn health_rid_matches_python_convention() { assert_eq!(HEALTH_CHECK_RID_PREFIX, "HEALTH_CHECK"); - let rid = new_health_check_rid(); + let rid = Rid::new_health_check(); // "HEALTH_CHECK_" + 32 hex chars assert!(rid.starts_with("HEALTH_CHECK_")); assert_eq!(rid.len(), "HEALTH_CHECK_".len() + 32); diff --git a/rust/sglang-server/src/lib.rs b/rust/sglang-server/src/lib.rs index 753c8d617..f733bb9dd 100644 --- a/rust/sglang-server/src/lib.rs +++ b/rust/sglang-server/src/lib.rs @@ -15,6 +15,7 @@ mod environ; mod error; mod fsm; mod ids; +mod message; mod utils; use pyo3::prelude::*; diff --git a/rust/sglang-server/src/message.rs b/rust/sglang-server/src/message.rs new file mode 100644 index 000000000..9b6e2754f --- /dev/null +++ b/rust/sglang-server/src/message.rs @@ -0,0 +1,18 @@ +//! Messages moved between stages via `flume` (zero-copy moves); variable-length +//! buffers are `bytes::Bytes`, so egress fan-out to detok shards is a refcount bump. +//! Grouped by flow direction: [`request`] (the `/generate` body fan-out, the +//! in-flight request bodies + scheduler ingress wire), [`egress`] +//! (the response back-channel + egress-ring frames and decoded chunk events), +//! [`finish_reason`] (the terminal reason a request ended, Python's +//! `FinishReasonDict`), [`sampling`] (sampling-params normalization, the Python +//! `SamplingParams` port), [`io_struct`] (the scheduler wire structs), [`types`] +//! (the shared wire-shape adapters both directions use). + +mod io_struct; +mod request; +mod sampling; +mod types; + +pub(crate) use request::GenerateRequest; +pub(crate) use sampling::{SamplingParams, SamplingParamsInput}; +pub(crate) use types::{OneOrMany, OneOrManyItem, TokenIds}; diff --git a/rust/sglang-server/src/message/io_struct.rs b/rust/sglang-server/src/message/io_struct.rs new file mode 100644 index 000000000..c3f0c7eb3 --- /dev/null +++ b/rust/sglang-server/src/message/io_struct.rs @@ -0,0 +1,173 @@ +//! The scheduler wire structs — the Rust mirror of the Python `io_struct` +//! messages this server sends (`python/sglang/srt/managers/io_struct.py`). +//! Each is a msgspec `array_like=True` struct, so **field order is wire order** +//! and `rmp_serde`'s default struct-as-array encoding reproduces it. + +use bytes::Bytes; +use serde::Serialize; + +use super::types::{Tagged, control_messages, wire_struct}; +use super::{GenerateRequest, SamplingParams, TokenIds}; +use crate::error::Error; + +wire_struct! { + /// The scheduler's `TokenizedGenerateReqInput`. Keep in lockstep with the + /// Python declaration: inserting a field anywhere but the end shifts every + /// later field on the wire. + pub(super) TokenizedGenerateReqInput<'a> { + input_text: Option<&'a str>, + /// Always nil: the ids ride the ring's columnar buffer, not msgpack. + input_ids: (), + input_embeds: (), + mm_inputs: (), + token_type_ids: (), + sampling_params: &'a SamplingParams, + return_logprob: bool, + logprob_start_len: i64, + top_logprobs_num: i64, + token_ids_logprob: Option<&'a TokenIds>, + stream: bool, + /// Not exposed by this server yet; the scheduler needs the slot filled. + return_sampling_mask: bool, + return_hidden_states: bool, + } +} + +// Owned-rid messages: these are held by a [`ControlRequest`] inside an owned +// `Request`, so they cannot borrow the rid that request owns. `pub(crate)` +// because that enum is crate-visible — their fields stay private, so only the +// constructors below can build one. +control_messages! { + /// The scheduler's `AbortReq`: stop generating for one rid. + AbortReq { + /// This server never aborts the whole queue — only the one rid. + abort_all: bool, + finished_reason: (), + abort_message: (), + } + + /// `/server_info`'s control request: a bare `BaseReq` with no extra fields. + GetInternalStateReq {} +} + +/// Borrow a request as its wire struct, resolving `Option` scalars to the wire +/// defaults Python's own fields carry. Borrowed, not owned: every field is a +/// reference into `req`, so an owning `From` would return references to a +/// dropped local. +/// +/// The rid comes from [`GenerateRequest::rid`] — the same value `submit` copied +/// into the owning `Request`, so the scheduler, the detok registration and +/// `meta_info.id` cannot disagree. +impl<'a> From<&'a GenerateRequest> for TokenizedGenerateReqInput<'a> { + fn from(req: &'a GenerateRequest) -> Self { + Self { + rid: &req.rid, + input_text: req.text.as_deref(), + input_ids: (), + input_embeds: (), + mm_inputs: (), + token_type_ids: (), + sampling_params: &req.sampling_params, + return_logprob: req.return_logprob, + logprob_start_len: req.logprob_start_len, + top_logprobs_num: req.top_logprobs_num, + token_ids_logprob: req.token_ids_logprob.as_ref(), + stream: req.stream, + return_sampling_mask: req.return_sampling_mask, + return_hidden_states: req.return_hidden_states, + } + } +} + +impl GetInternalStateReq { + pub fn new(rid: String) -> Self { + Self { rid } + } +} + +impl AbortReq { + pub fn new(rid: String, abort_all: bool) -> Self { + Self { + rid, + abort_all, + finished_reason: (), + abort_message: (), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn abort_req_msgpack_shape() { + let b = AbortReq::new("12345".into(), false).encode().unwrap(); + let val = rmpv::decode::read_value(&mut &b[..]).unwrap(); + let arr = val.as_array().expect("array"); + assert_eq!( + arr.len(), + 6, + "AbortReq = [tag, rid, http_ipc, abort_all, finished_reason, abort_message]" + ); + assert_eq!(arr[0].as_str(), Some("AbortReq")); + assert_eq!(arr[1].as_str(), Some("12345")); + assert!(arr[2].is_nil()); + assert_eq!(arr[3].as_bool(), Some(false)); + assert!(arr[4].is_nil()); + assert!(arr[5].is_nil()); + } + + /// The header must be positionally aligned: `input_embeds` (idx 5) / + /// `token_type_ids` (idx 7) present as nil so `sampling_params` lands at idx 8 and + /// the array reaches msgspec's min length. Regression guard for that decode failure. + #[test] + fn to_header_msgpack_is_positionally_aligned() { + let req = GenerateRequest { + rid: "r1".into(), + text: Some("hi".into()), + input_ids: Some(vec![1, 2, 3]), + sampling_params: SamplingParams { + max_new_tokens: Some(5), + ..Default::default() + }, + return_logprob: true, + logprob_start_len: -1, + top_logprobs_num: 3, + return_hidden_states: true, + stream: true, + ..Default::default() + }; + let bytes = TokenizedGenerateReqInput::from(&req).encode().unwrap(); + let val = rmpv::decode::read_value(&mut &bytes[..]).unwrap(); + let arr = val.as_array().expect("array"); + // msgspec requires >= 14 (through `stream`); we emit 16. + assert!( + arr.len() >= 14, + "header must have >=14 elements, got {}", + arr.len() + ); + assert_eq!(arr[0].as_str(), Some("TokenizedGenerateReqInput")); + assert_eq!(arr[1].as_str(), Some("r1")); + assert!(arr[5].is_nil(), "idx 5 must be input_embeds (nil)"); + assert!(arr[7].is_nil(), "idx 7 must be token_type_ids (nil)"); + // An ARRAY, not a map: Python's `SamplingParams` is + // `msgspec.Struct(array_like=True)`, so it decodes positionally. + assert!(arr[8].is_array(), "sampling_params must land at idx 8"); + assert_eq!(arr[9].as_bool(), Some(true), "return_logprob at idx 9"); + assert_eq!(arr[11].as_u64(), Some(3), "top_logprobs_num at idx 11"); + assert_eq!(arr[13].as_bool(), Some(true), "stream at idx 13"); + // idx 14 is `return_sampling_mask` (never client-set); a shift here would + // silently flip the wrong scheduler field. + assert_eq!( + arr[14].as_bool(), + Some(false), + "return_sampling_mask at idx 14" + ); + assert_eq!( + arr[15].as_bool(), + Some(true), + "return_hidden_states at idx 15" + ); + } +} diff --git a/rust/sglang-server/src/message/request.rs b/rust/sglang-server/src/message/request.rs new file mode 100644 index 000000000..9b2c5e3d9 --- /dev/null +++ b/rust/sglang-server/src/message/request.rs @@ -0,0 +1,820 @@ +//! The `/generate` request path: the HTTP body and its per-request fan-out +//! ([`GenerateBody`] → [`GenerateRequest`]s), the variant bodies, and the +//! scheduler ingress encodings (`TokenizedGenerateReqInput` header, +//! control/abort, `IngressMsg`). + +use std::collections::HashSet; +use std::sync::LazyLock; + +use bytes::Bytes; +use itertools::izip; +use serde::Deserialize; + +use super::io_struct::{ControlRequest, TokenizedGenerateReqInput}; +use super::{OneOrMany, OneOrManyItem, SamplingParams, SamplingParamsInput, TokenIds}; +use crate::environ::env_u64; +use crate::error::Error; +use crate::ids::Rid; + +/// Hard cap on how many scheduler requests one `/generate` HTTP call may expand +/// into. Every column below is allocated per item before anything is dispatched, +/// so this bounds the work — and the resident memory — a single call can ask for. +/// +/// NOT a concurrency limit: it is a pure function of the body being parsed, so +/// separate HTTP calls never interact with it. +/// +/// Read once from `SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ` (registered in +/// `python/sglang/srt/environ.py`, which owns the default). Memoized because the +/// value is process-static — Python sets it before launching this server — and a +/// per-request `env::var` would take a lock on the hot path for a constant. +static MAX_BATCH_REQS_PER_HTTP_REQ: LazyLock = + LazyLock::new(|| env_u64("SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ", 4096) as usize); + +/// Hard cap on the total bytes a broadcast value may clone into the batch (see +/// the `One` arms of the fan-out). +const MAX_BROADCAST_CLONE_BYTES: usize = 64 << 20; + +/// Live heap per byte of serialized JSON. Measured across shapes at 1.0–7.0× +/// (`serde_json::Value` pays for enum tags, `String` headers and map nodes that +/// the wire form does not); 8 is the ceiling of that range, not a worst case. +const JSON_TO_HEAP_FACTOR: usize = 8; + +/// The `/generate` wire body before batch splitting: `text`/`input_ids`/`sampling_params` +/// each scalar-or-list, fanned into per-request [`GenerateRequest`]s by +/// [`into_requests`](GenerateBody::into_requests). +/// +/// Unknown keys are IGNORED, matching Python: FastAPI builds `GenerateReqInput` +/// as a pydantic dataclass, which drops extras. `deny_unknown_fields` here turned +/// every `GenerateReqInput` field this server has not ported — `priority`, +/// `extra_key`, `session_id`, `session_params`, `return_sampling_mask`, +/// `custom_logit_processor`, and ~40 more — into a 400, so a client that worked +/// against the Python server broke against this one. The cost of dropping it is +/// that a typo (`temperature`) is silently ignored rather than reported; that is +/// the same trade Python already makes. +#[derive(Debug, Clone, Default, Deserialize)] +pub struct GenerateBody { + /// Optional client-supplied request id(s): a single string (a batch fans it + /// out as `{rid}_{i}`, mirroring Python `_normalize_batch`) or one per item. + #[serde(default)] + pub rid: Option>, + #[serde(default)] + pub text: Option>, + #[serde(default)] + pub input_ids: Option>, + #[serde(default)] + pub stream: bool, + /// One params object (broadcast) or a list of them (per item); see + /// [`SamplingParamsInput`]. + #[serde(default)] + pub sampling_params: Option, + /// Logprob / hidden-state options: a scalar broadcasts to every prompt, a + /// list is per-prompt (Python `_normalize_logprob_params`). + #[serde(default)] + pub return_logprob: Option>, + #[serde(default)] + pub logprob_start_len: Option>, + #[serde(default)] + pub top_logprobs_num: Option>, + /// Token ids to report logprobs for: one list (broadcast to every prompt) or + /// one list per prompt, mirroring Python's + /// `Union[List[int], List[List[int]]]` fan-out in `_normalize_batch`. + #[serde(default)] + pub token_ids_logprob: Option>, + #[serde(default)] + pub return_hidden_states: Option>, + /// Scalar-only in Python too (`return_text_in_logprobs: bool`). + #[serde(default)] + pub return_text_in_logprobs: Option, +} + +impl GenerateBody { + /// 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 + /// `GenerateReqInput.normalize_batch_and_arguments`; an invalid/inconsistent + /// batch is [`Error::Validation`], which the handler surfaces with the + /// variant's own status (400). + pub fn into_requests(self) -> Result<(Vec, bool), Error> { + let GenerateBody { + rid, + text, + input_ids, + stream, + sampling_params, + return_logprob, + logprob_start_len, + top_logprobs_num, + token_ids_logprob, + return_hidden_states, + return_text_in_logprobs, + // Unported `GenerateReqInput` fields land here and are dropped, as they + // are on the Python path. + .. + } = self; + + // Cap the batch BEFORE the columns below allocate anything. Reading the + // declared length off the input costs nothing; the previous placement (after + // the match) had already allocated ~1.7 GiB for a 114 MiB body, most of it + // the `vec![None; n]` twin column. + let declared_n = match (&text, &input_ids) { + (Some(OneOrMany::Many(v)), None) => v.len(), + (None, Some(OneOrMany::Many(v))) => v.len(), + _ => 1, + }; + if declared_n > *MAX_BATCH_REQS_PER_HTTP_REQ { + return Err(Error::Validation(format!( + "batch size {declared_n} exceeds the maximum of {}", + *MAX_BATCH_REQS_PER_HTTP_REQ + ))); + } + + // Per-item (text, input_ids) columns + whether the input used list form. + type Columns = (Vec>, Vec>, bool); + // Exactly one of text / input_ids (Python `_validate_inputs`), and no + // empty id list (Python `_determine_batch_size`). + let (texts, id_lists, is_batch): Columns = match (text, input_ids) { + (Some(_), Some(_)) => { + return Err(Error::Validation( + "provide either `text` or `input_ids`, not both".into(), + )); + } + (None, None) => { + return Err(Error::Validation( + "either `text` or `input_ids` must be provided".into(), + )); + } + (Some(OneOrMany::One(s)), None) => (vec![Some(s)], vec![None], false), + (Some(OneOrMany::Many(v)), None) => { + let n = v.len(); + (v.into_iter().map(Some).collect(), vec![None; n], true) + } + // `[]` parses as `One(vec![])` (one prompt with no ids), so the + // `n == 0` guard below never sees it — reject it here, as Python's + // `_determine_batch_size` does. + (None, Some(OneOrMany::One(x))) => { + if x.is_empty() { + return Err(Error::Validation("input_ids cannot be empty".into())); + } + (vec![None], vec![Some(x)], false) + } + (None, Some(OneOrMany::Many(vv))) => { + if vv.iter().any(|ids| ids.is_empty()) { + return Err(Error::Validation( + "input_ids cannot be empty for any prompt in the batch".into(), + )); + } + let n = vv.len(); + (vec![None; n], vv.into_iter().map(Some).collect(), true) + } + }; + let n = texts.len(); + if n == 0 { + return Err(Error::Validation( + "batch must contain at least one item".into(), + )); + } + + // A list is per-item; a single object broadcasts to every item. + let sps: Vec = match sampling_params { + None => vec![SamplingParams::default(); n], + Some(SamplingParamsInput::Many(v)) => { + if v.len() != n { + return Err(Error::Validation(format!( + "sampling_params list length {} does not match batch size {n}", + v.len() + ))); + } + v + } + Some(SamplingParamsInput::One(sp)) => { + // Broadcasting deep-clones the client's params once per prompt, + // heap and all — `stop`, `logit_bias` and `custom_params` (arbitrary + // JSON) are still unnormalized client data here. The blow-up is + // quadratic in the body: ~1 MB of `custom_params` broadcast to 200k + // prompts is ~200 GB of clones, and a Rust allocation failure calls + // `abort()`, which is uncatchable and takes the scheduler process + // with it. Bound the product, not just `n`. + // `n == 1` is not a broadcast, so skip the sizing entirely: measuring + // it means serializing the client's whole `custom_params` to a + // throwaway `String` on every single request. The callee's own + // `n > 1` guard cannot prevent that — the cost is in the argument. + if n > 1 { + // Serialized bytes are NOT the clone cost: measured, 63.7 MiB of + // JSON became ~1008 MiB of live heap once parsed into `Value` + // nodes, `String`s and map entries. Scale by that measured factor + // so the budget bounds memory rather than wire size. + let per_clone = serde_json::to_string(&*sp) + .map_or(0, |s| s.len()) + .saturating_mul(JSON_TO_HEAP_FACTOR); + check_broadcast_budget(per_clone, n, "sampling_params")?; + } + vec![*sp; n] + } + }; + + // rid: absent → mint one uuid per item here, so every request carries its + // final rid from this point on; a single string fans out as `{rid}_{i}` + // for a batch (Python `_normalize_batch`); a list is per-item. + // + // Every CLIENT-supplied rid goes through `Rid::from_client`, which appends a + // uniquifier so two concurrent requests sharing an rid cannot collide on the + // detok table. `client_facing` strips it back off for `meta_info.id`, so the + // client sees exactly what it sent. Minted rids (`Rid::default`) are already + // unique and are left bare. + let rids: Vec = match rid { + None => (0..n).map(|_| Rid::default()).collect(), + Some(OneOrMany::One(r)) if !is_batch => vec![Rid::from_client(&r)], + Some(OneOrMany::One(r)) => { + check_broadcast_budget(r.len(), n, "rid")?; + // Uniquify AFTER the `_{i}` split, so the split index stays part of + // the rid the client gets back. + (0..n) + .map(|i| Rid::from_client(&format!("{r}_{i}"))) + .collect() + } + Some(OneOrMany::Many(v)) => { + if !is_batch || v.len() != n { + return Err(Error::Validation(format!( + "rid list length {} does not match batch size {n}", + v.len() + ))); + } + // Python `_validate_rid_uniqueness`. `from_client` below would make + // even these unique, so this is parity rather than safety: Python + // 400s a request that names one id twice, and echoing the same + // `meta_info.id` on two entries of one batch response is useless to + // the client regardless. Checked on the RAW strings, before the + // uniquifier hides the duplication. + { + let mut seen = HashSet::with_capacity(v.len()); + let duplicates: Vec<&String> = v.iter().filter(|r| !seen.insert(*r)).collect(); + if !duplicates.is_empty() { + return Err(Error::Validation(format!( + "duplicate request IDs detected within the request: {duplicates:?}" + ))); + } + } + v.iter().map(|r| Rid::from_client(r)).collect() + } + }; + + // Fans out exactly like the scalar options: one list broadcasts, a list of + // lists is per item (Python `_normalize_batch`'s nested branch). Empties + // are collapsed per item below, not here. + let tid_logprobs = fan_out(token_ids_logprob, n, "token_ids_logprob")?; + + // Each logprob/hidden opt: absent → None for every item, a scalar + // broadcasts, a list is per-item (Python `normalize_param`, plus a length + // check Python lacks — it would `IndexError` later instead). + let return_logprobs = fan_out(return_logprob, n, "return_logprob")?; + let logprob_start_lens = fan_out(logprob_start_len, n, "logprob_start_len")?; + let top_logprobs_nums = fan_out(top_logprobs_num, n, "top_logprobs_num")?; + let return_hidden = fan_out(return_hidden_states, n, "return_hidden_states")?; + + // Every column above is exactly `n` long, so zip them by value: each + // request takes ownership of its cell, with no indexing or bounds checks. + let requests = izip!( + rids, + texts, + id_lists, + sps, + return_logprobs, + logprob_start_lens, + top_logprobs_nums, + tid_logprobs, + return_hidden, + ) + .map( + |( + rid, + text, + input_ids, + sampling_params, + return_logprob, + logprob_start_len, + top_logprobs_num, + token_ids_logprob, + return_hidden_states, + )| GenerateRequest { + rid, + text, + input_ids, + sampling_params, + stream, + // Python `GenerateReqInput` defaults. + return_logprob: return_logprob.unwrap_or(false), + logprob_start_len: logprob_start_len.unwrap_or(-1), + top_logprobs_num: top_logprobs_num.unwrap_or(0), + // `Some` here means "these ids were requested", so an empty list + // collapses to None. + token_ids_logprob: token_ids_logprob.filter(|ids| !ids.is_empty()), + return_sampling_mask: false, // TODO: port Python's `return_sampling_mask` + return_hidden_states: return_hidden_states.unwrap_or(false), + return_text_in_logprobs, + }, + ) + .collect(); + Ok((requests, is_batch)) + } +} + +/// Request variant — selects the ingress branch, scheduler wire message, and +/// egress shape. Each owns its body, so generate/control fields stay type-separate. +#[derive(Debug)] +pub enum RequestKind { + /// `/generate`: tokenize (if needed) then push a `TokenizedGenerateReqInput`. + Generate(Box), + /// A control endpoint (e.g. `/server_info`, `/health`): no tokenization, and + /// the egress is a single non-streamed JSON result. + Control(Box), +} + +/// A single in-flight `/generate` request (per-item from +/// [`GenerateBody::into_requests`]), +/// serialized to the scheduler wire once tokenized (see `to_header_msgpack`). Not a +/// wire type — built by `into_requests`/handlers, never (de)serialized; `input_ids` is +/// client-supplied or filled by the Tokenizer stage. +#[derive(Debug, Default)] +pub struct GenerateRequest { + /// This item's final rid: the client's (normalized per item by `into_requests`) or a + /// uuid minted there when none was sent. A [`Rid`], not a `String`: the wire + /// forms stay textual (`GenerateBody` on the way in, `TokenizedGenerateReqInput` + /// on the way out) but every in-process carrier names the type. + /// + /// Duplicates *within* one request are rejected by `into_requests` (Python + /// `_validate_rid_uniqueness`). A collision with a *concurrent* request's rid + /// cannot arise: [`Rid::from_client`] appends a uniquifier to every + /// client-supplied rid, so this value is unique for the process's lifetime and + /// only [`client_facing`](Rid::client_facing) is ever shown back. + /// + /// This diverges from Python, which 400s the second request ("Duplicate request + /// ID detected"). Serving both is the friendlier answer and strictly safer — + /// what the rejection protected against was one request evicting the other's + /// detok sink, which is now unrepresentable. + pub rid: Rid, + pub text: Option, + /// Client-supplied token ids, or filled by the Tokenizer stage. + pub input_ids: Option, + /// Sampling params (defaults when the client sent none, as in Python); + /// normalized + verified at ingress, then serialized into the header. + pub sampling_params: SamplingParams, + /// Whether the client asked for SSE streaming. + pub stream: bool, + /// Logprob / hidden-state options. This path bypasses the Python + /// `TokenizerManager`, so `into_requests` replicates its scalar + /// normalization. Resolved to concrete values THERE rather than at the wire + /// boundary: an `Option` surviving past construction invites two call sites + /// to disagree about what absent means, and only the wire knew the answer. + /// The defaults are `GenerateReqInput`'s own. + pub return_logprob: bool, + pub logprob_start_len: i64, + pub top_logprobs_num: i64, + /// This request's `token_ids_logprob` ids, fanned out by `into_requests` and + /// collapsed to `None` when empty (the scheduler branches on `is not None`). + pub token_ids_logprob: Option, + pub return_sampling_mask: bool, + pub return_hidden_states: bool, + /// Decode logprob token ids to text in each `[logprob, token_id, text]` tuple + /// (default leaves the text slot null). Deliberately NOT in the scheduler + /// header — Python's `TokenizedGenerateReqInput` has no such field either; + /// it is consumed on the way out, by `register_detok` → `DetokMsg::Register` + /// → the shard's `decode_logprob_texts`. + pub return_text_in_logprobs: Option, +} + +impl GenerateRequest { + /// True when the client already supplied token ids → skip tokenization. + pub fn already_tokenized(&self) -> bool { + self.input_ids.as_ref().is_some_and(|v| !v.is_empty()) + } + + /// Multimodal detection hook. Deferred (Encoder stubbed): always false until mm + /// fields are wired in. + #[allow(dead_code)] + pub fn has_multimodal(&self) -> bool { + false + } + + pub fn encode_header(&self) -> Result { + TokenizedGenerateReqInput::from(self).encode() + } + + /// `input_ids` widened to raw little-endian int64 bytes (the scheduler's + /// `array("q")` columnar cell — rides the ingress ring outside msgpack). Empty + /// when not tokenized. + pub fn encode_data_buf(&self) -> Bytes { + let ids = self.input_ids.as_deref().unwrap_or(&[]); + let mut buf = Vec::with_capacity(ids.len() * 8); + for &id in ids { + buf.extend_from_slice(&(id as i64).to_le_bytes()); + } + Bytes::from(buf) + } +} + +/// Fan one scalar-or-list option out to `n` per-item values: absent → `None` +/// each, a scalar broadcasts, a list must match the batch size. +/// Bytes a broadcast value costs per clone. Only the heap matters — the inline +/// part is bounded by the type. +trait HeapBytes { + fn heap_bytes(&self) -> usize; +} +impl HeapBytes for bool { + fn heap_bytes(&self) -> usize { + 0 + } +} +impl HeapBytes for i64 { + fn heap_bytes(&self) -> usize { + 0 + } +} +impl HeapBytes for String { + fn heap_bytes(&self) -> usize { + self.len() + } +} +impl HeapBytes for TokenIds { + fn heap_bytes(&self) -> usize { + self.len() * std::mem::size_of::() + } +} + +/// Reject a broadcast whose clones would exceed [`MAX_BROADCAST_CLONE_BYTES`]. +fn check_broadcast_budget(per_clone: usize, n: usize, name: &str) -> Result<(), Error> { + // `n == 1` is not a broadcast — there is one value and one prompt, so nothing + // is duplicated. Charging it here rejected ordinary single requests with a + // message about a batch they never sent. + if n > 1 && per_clone.saturating_mul(n) > MAX_BROADCAST_CLONE_BYTES { + return Err(Error::Validation(format!( + "{name} ({per_clone} bytes) broadcast to {n} prompts would allocate more \ + than the {MAX_BROADCAST_CLONE_BYTES}-byte limit; send a shorter {name} \ + or a smaller batch" + ))); + } + Ok(()) +} + +fn fan_out( + value: Option>, + n: usize, + name: &str, +) -> Result>, Error> { + match value { + None => Ok(vec![None; n]), + Some(OneOrMany::One(v)) => { + // Same budget as the `sampling_params` broadcast: `vec![Some(v); n]` + // deep-clones client data once per prompt, so a 16 MiB + // `token_ids_logprob` fanned to 4096 prompts is ~64 GiB — an + // allocation failure, which `abort()`s the scheduler process. + check_broadcast_budget(v.heap_bytes(), n, name)?; + Ok(vec![Some(v); n]) + } + Some(OneOrMany::Many(v)) => { + if v.len() != n { + return Err(Error::Validation(format!( + "{name} list length {} does not match batch size {n}", + v.len() + ))); + } + Ok(v.into_iter().map(Some).collect()) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn requests(body: &str) -> Result<(Vec, bool), Error> { + serde_json::from_str::(body) + .unwrap() + .into_requests() + } + + /// Scalar `text` → one item, not a batch (response stays a single object). + #[test] + fn scalar_text_is_single() { + let (ps, is_batch) = requests(r#"{"text": "hi"}"#).unwrap(); + assert!(!is_batch); + assert_eq!(ps.len(), 1); + assert_eq!(ps[0].text.as_deref(), Some("hi")); + } + + /// List `text` → batch (even length 1); each prompt becomes its own payload. + #[test] + fn list_text_is_batch() { + let (ps, is_batch) = requests(r#"{"text": ["a", "b"]}"#).unwrap(); + assert!(is_batch); + assert_eq!(ps.len(), 2); + assert_eq!(ps[0].text.as_deref(), Some("a")); + assert_eq!(ps[1].text.as_deref(), Some("b")); + + let (ps, is_batch) = requests(r#"{"text": ["only"]}"#).unwrap(); + assert!(is_batch, "single-element list is still a batch"); + assert_eq!(ps.len(), 1); + } + + /// Scalar `sampling_params` broadcasts to every item; a list maps per item. + #[test] + fn sampling_params_broadcast_and_per_item() { + let (ps, _) = + requests(r#"{"text": ["a", "b"], "sampling_params": {"temperature": 0.5}}"#).unwrap(); + assert_eq!(ps[0].sampling_params, ps[1].sampling_params); + assert_eq!(ps[0].sampling_params.temperature, 0.5); + + let (ps, _) = requests( + r#"{"text": ["a", "b"], "sampling_params": [{"temperature": 0.1}, {"temperature": 0.9}]}"#, + ) + .unwrap(); + assert_ne!(ps[0].sampling_params, ps[1].sampling_params); + } + + /// A per-item `sampling_params` list whose length ≠ batch size is a 400. + #[test] + fn sampling_params_length_mismatch_errors() { + let err = requests(r#"{"text": ["a", "b"], "sampling_params": [{}]}"#).unwrap_err(); + assert!(err.to_string().contains("length"), "{err}"); + } + + /// `input_ids` batch (list of lists) fans out; scalar (list of ints) is single. + #[test] + fn input_ids_scalar_vs_batch() { + let (ps, is_batch) = requests(r#"{"input_ids": [1, 2, 3]}"#).unwrap(); + assert!(!is_batch); + assert_eq!(ps[0].input_ids, Some(vec![1, 2, 3])); + + let (ps, is_batch) = requests(r#"{"input_ids": [[1, 2], [3]]}"#).unwrap(); + assert!(is_batch); + assert_eq!(ps.len(), 2); + assert_eq!(ps[1].input_ids, Some(vec![3])); + } + + /// Both / neither of text+input_ids is a 400. + #[test] + fn split_validates_inputs() { + assert!(requests(r#"{"text": "a", "input_ids": [1]}"#).is_err()); + assert!(requests(r#"{"stream": true}"#).is_err()); + // Parallel sampling is rejected where Python reads it — in the params, + // at normalization (the ingress step), not here. + let (mut ps, _) = requests(r#"{"text": "a", "sampling_params": {"n": 2}}"#).unwrap(); + assert!(ps[0].sampling_params.normalize(false, None).is_err()); + } + + /// Unported `GenerateReqInput` fields are IGNORED, not rejected. + /// + /// These are all real fields on Python's `GenerateReqInput` that this server + /// has not ported. `deny_unknown_fields` turned every one of them into a 400, + /// so a client that worked against the Python server broke here — and the + /// wire-compat fields (`lora_path`, `image_data`, `return_routed_experts`) had + /// to be declared and dropped by hand just to let `bench_serving` through. + /// FastAPI's pydantic dataclass drops extras, so ignoring them is the parity + /// behavior; a typo being silently ignored is the same trade Python makes. + #[test] + fn unported_generate_req_input_fields_are_ignored() { + for field in [ + r#""priority": 3"#, + r#""extra_key": "k""#, + r#""session_id": "s""#, + r#""session_params": {"a": 1}"#, + r#""return_sampling_mask": true"#, + r#""custom_logit_processor": "cls""#, + r#""lora_path": "adapter""#, + r#""image_data": "base64""#, + r#""return_routed_experts": true"#, + r#""bootstrap_host": "h""#, + // Python has no top-level `n` either, and ignores it just the same. + r#""n": 1"#, + r#""totally_made_up": 1"#, + ] { + let body = format!(r#"{{"text": "hi", {field}}}"#); + let (ps, _) = requests(&body) + .unwrap_or_else(|e| panic!("{field} must be ignored, not rejected: {e}")); + assert_eq!(ps.len(), 1, "{field}"); + assert_eq!(ps[0].text.as_deref(), Some("hi"), "{field}"); + } + } + + /// Client-supplied rid semantics mirror Python's `_normalize_batch`: a + /// single string passes through for a single request, fans out as + /// `{rid}_{i}` for a batch, and a list must match the batch length. An + /// absent rid is minted here, one uuid per item. + /// + /// Asserted on `client_facing()`, which is what `meta_info.id` echoes: the + /// internal rid additionally carries the `from_client` uniquifier, and that + /// suffix must never be visible in the parity-defined shape. + #[test] + fn split_rid_matches_python_normalize() { + let (ps, _) = requests(r#"{"text": "a", "rid": "r"}"#).unwrap(); + assert_eq!(ps[0].rid.client_facing(), "r"); + + let (ps, _) = requests(r#"{"text": ["a", "b"], "rid": "base"}"#).unwrap(); + assert_eq!(ps[0].rid.client_facing(), "base_0"); + assert_eq!(ps[1].rid.client_facing(), "base_1"); + + let (ps, _) = requests(r#"{"text": ["a", "b"], "rid": ["x", "y"]}"#).unwrap(); + assert_eq!(ps[0].rid.client_facing(), "x"); + assert_eq!(ps[1].rid.client_facing(), "y"); + + let (ps, _) = requests(r#"{"text": ["a", "b"]}"#).unwrap(); + // Absent → `into_requests` mints one uuid per item, all distinct. + assert_eq!(ps[0].rid.len(), 32); + assert_ne!(ps[0].rid, ps[1].rid); + + assert!( + requests(r#"{"text": ["a", "b"], "rid": ["x"]}"#).is_err(), + "rid list length must match batch size" + ); + assert!( + requests(r#"{"text": "a", "rid": ["x"]}"#).is_err(), + "rid list with a single (non-batch) prompt is rejected" + ); + } + + /// The native `bench_serving` payload (a `GenerateReqInput` superset) parses: + /// its `lora_path`/`return_routed_experts`/`image_data` are accepted-but-ignored, + /// so `split` succeeds and drops them while the real fields survive. + #[test] + fn accepts_bench_serving_payload() { + let (ps, is_batch) = requests( + r#"{"text": "hi", "sampling_params": {"max_new_tokens": 8}, + "stream": true, "lora_path": null, "return_logprob": false, + "return_routed_experts": false, "logprob_start_len": -1, + "image_data": null}"#, + ) + .unwrap(); + assert!(!is_batch); + assert_eq!(ps.len(), 1); + assert_eq!(ps[0].text.as_deref(), Some("hi")); + assert!(ps[0].stream); + } + + /// The body limit is disabled, so an unbounded batch turns a small body into an + /// unbounded allocation. Worse, broadcasting `sampling_params` deep-clones the + /// client's `custom_params`/`logit_bias`/`stop` once per prompt, so the blow-up + /// is quadratic in the body — and a Rust allocation failure `abort()`s the + /// scheduler process rather than raising. Both the count and the product are + /// capped before any column is built. + #[test] + fn oversized_batches_are_rejected_before_allocating() { + let texts: Vec = (0..*MAX_BATCH_REQS_PER_HTTP_REQ + 1) + .map(|i| i.to_string()) + .collect(); + let body = serde_json::json!({ "text": texts }).to_string(); + let err = requests(&body).unwrap_err().to_string(); + assert!(err.contains("exceeds the maximum"), "{err}"); + + // At the cap it is accepted. + let texts: Vec = (0..*MAX_BATCH_REQS_PER_HTTP_REQ) + .map(|i| i.to_string()) + .collect(); + let (reqs, _) = requests(&serde_json::json!({ "text": texts }).to_string()).unwrap(); + assert_eq!(reqs.len(), *MAX_BATCH_REQS_PER_HTTP_REQ); + + // A small batch with a huge broadcast `custom_params` is the quadratic case: + // few items, but each clone carries the whole blob. The item count is a + // literal because this half asserts the BYTE budget, not the item cap — + // it therefore assumes the default `SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ`, since + // a cap below 200 would trip the item check first and report that instead. + let blob = "x".repeat(1 << 20); // 1 MiB + let body = serde_json::json!({ + "text": vec!["hi"; 200], + "sampling_params": { "custom_params": { "k": blob } }, + }) + .to_string(); + let err = requests(&body).unwrap_err().to_string(); + assert!(err.contains("would allocate more than"), "{err}"); + } + + /// `token_ids_logprob` mirrors Python `_normalize_batch`'s nested-structure + /// branch: a flat list broadcasts to every prompt, a list of lists is + /// per-prompt. Regression — the whole value used to be cloned to every item. + #[test] + fn token_ids_logprob_broadcasts_flat_and_splits_nested() { + let (ps, _) = requests(r#"{"text": ["a", "b"], "token_ids_logprob": [1, 2]}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, Some(vec![1, 2])); + assert_eq!(ps[1].token_ids_logprob, Some(vec![1, 2])); + + let (ps, _) = + requests(r#"{"text": ["a", "b"], "token_ids_logprob": [[1], [2, 3]]}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, Some(vec![1])); + assert_eq!(ps[1].token_ids_logprob, Some(vec![2, 3])); + + let err = requests(r#"{"text": ["a", "b"], "token_ids_logprob": [[1]]}"#).unwrap_err(); + assert!( + err.to_string().contains("does not match batch size"), + "{err}" + ); + + let (ps, _) = requests(r#"{"text": ["a", "b"]}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, None); + } + + /// An empty `token_ids_logprob` means "none requested" and must reach the + /// scheduler as None, whose guards are `x is not None` — `Some([])` enters the + /// token-ids-logprob path and computes nothing. The collapse is per item, so it + /// covers every shape: Python only collapses the outer value + /// (`if not self.token_ids_logprob`, io_struct.py:439,612) and passes inner + /// empties through its nested branch verbatim. + #[test] + fn empty_token_ids_logprob_collapses_to_none() { + let (ps, _) = requests(r#"{"text": "a", "token_ids_logprob": []}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, None); + + let (ps, _) = requests(r#"{"text": ["a", "b"], "token_ids_logprob": []}"#).unwrap(); + assert!(ps.iter().all(|p| p.token_ids_logprob.is_none())); + + // Nested, every item empty — Python would ship four `[]`s here. + let (ps, _) = + requests(r#"{"text": ["a", "b", "c", "d"], "token_ids_logprob": [[], [], [], []]}"#) + .unwrap(); + assert!(ps.iter().all(|p| p.token_ids_logprob.is_none())); + + // Nested and mixed: only the empty cell collapses. + let (ps, _) = requests(r#"{"text": ["a", "b"], "token_ids_logprob": [[], [7]]}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, None); + assert_eq!(ps[1].token_ids_logprob, Some(vec![7])); + + // A non-empty list is untouched. + let (ps, _) = requests(r#"{"text": "a", "token_ids_logprob": [7]}"#).unwrap(); + assert_eq!(ps[0].token_ids_logprob, Some(vec![7])); + } + + /// The logprob/hidden options take Python's batch form too + /// (`Union[List[T], T]`): a scalar broadcasts, a list is per-prompt. + #[test] + fn logprob_options_broadcast_scalar_and_split_list() { + let (ps, _) = + requests(r#"{"text": ["a", "b"], "return_logprob": true, "top_logprobs_num": 3}"#) + .unwrap(); + assert!(ps[0].return_logprob); + assert_eq!(ps[1].top_logprobs_num, 3); + + let (ps, _) = requests( + r#"{"text": ["a", "b"], "return_logprob": [true, false], + "logprob_start_len": [0, 2], "return_hidden_states": [false, true]}"#, + ) + .unwrap(); + assert!(ps[0].return_logprob); + assert!(!ps[1].return_logprob); + assert_eq!(ps[0].logprob_start_len, 0); + assert_eq!(ps[1].logprob_start_len, 2); + assert!(ps[1].return_hidden_states); + + let err = requests(r#"{"text": ["a", "b"], "return_logprob": [true]}"#).unwrap_err(); + assert!( + err.to_string().contains("does not match batch size"), + "{err}" + ); + } + + /// `{"input_ids": []}` parses as one prompt with no ids, so the batch-size + /// guard misses it; Python's `_determine_batch_size` raises "input_ids cannot + /// be empty." Regression — it used to reach the tokenizer with no text. + #[test] + fn empty_input_ids_is_rejected() { + let err = requests(r#"{"input_ids": []}"#).unwrap_err(); + assert!( + err.to_string().contains("input_ids cannot be empty"), + "{err}" + ); + + let err = requests(r#"{"input_ids": [[1, 2], []]}"#).unwrap_err(); + assert!(err.to_string().contains("cannot be empty"), "{err}"); + + assert!(requests(r#"{"input_ids": [1, 2]}"#).is_ok()); + assert!(requests(r#"{"input_ids": [[1], [2]]}"#).is_ok()); + } + + /// Two items in one request cannot share an rid. Mirrors Python + /// `_validate_rid_uniqueness` — and it must be checked on the RAW strings, + /// because `Rid::from_client` would otherwise make the duplicates distinct and + /// the client would get two response entries carrying the same `meta_info.id`. + #[test] + fn duplicate_rids_within_one_request_are_rejected() { + let err = requests(r#"{"text": ["a", "b"], "rid": ["x", "x"]}"#).unwrap_err(); + assert!(err.to_string().contains("duplicate request IDs"), "{err}"); + + assert!(requests(r#"{"text": ["a", "b"], "rid": ["x", "y"]}"#).is_ok()); + let (ps, _) = requests(r#"{"text": ["a", "b"], "rid": "x"}"#).unwrap(); + assert_eq!(ps[0].rid.client_facing(), "x_0"); + assert_eq!(ps[1].rid.client_facing(), "x_1"); + } + + /// The collision this whole scheme exists to prevent: two CONCURRENT requests + /// naming the same rid. They must end up with different internal `Rid`s — the + /// detok table is keyed on it, and `Register` is an insert-overwrite, so equal + /// rids would evict the first client's sink and deliver its remaining chunks to + /// the second's connection. Both still see their own rid echoed back. + #[test] + fn concurrent_requests_sharing_an_rid_get_distinct_internal_rids() { + let (a, _) = requests(r#"{"text": "a", "rid": "same"}"#).unwrap(); + let (b, _) = requests(r#"{"text": "b", "rid": "same"}"#).unwrap(); + assert_ne!( + a[0].rid, b[0].rid, + "a shared client rid must not become a shared internal rid" + ); + assert_eq!(a[0].rid.client_facing(), "same"); + assert_eq!(b[0].rid.client_facing(), "same"); + } +} diff --git a/rust/sglang-server/src/message/sampling.rs b/rust/sglang-server/src/message/sampling.rs new file mode 100644 index 000000000..309dfbafb --- /dev/null +++ b/rust/sglang-server/src/message/sampling.rs @@ -0,0 +1,80 @@ +//! [`SamplingParams`] — the typed Rust port of Python `SamplingParams` +//! (python/sglang/srt/sampling/sampling_params.py): every field, plus its +//! `__post_init__` → `normalize` → `verify` pipeline (run in that order, as +//! `TokenizerManager._create_tokenized_object` does). +//! +//! The embedded Rust server replaces the Python `TokenizerManager`, which is the +//! only place those three run on the normal (zmq) path. Running them here, in the +//! ingress `Normalizing` FSM step, keeps the per-request CPU (notably the +//! stop-string work) off the scheduler's latency-critical loop. We set +//! `is_normalized=true` on the wire so the scheduler's `__post_init__` and +//! `normalize` early-return; its `verify` is likewise skipped (we did it here). +//! +//! KEEP IN SYNC with `sampling_params.py`: the field list, defaults and ranges +//! below mirror that file, and the struct is serialized by field name into the +//! `TokenizedGenerateReqInput` header, so a renamed/added Python field must be +//! mirrored here (an unknown key would be silently dropped by msgspec). +//! +//! Two deliberate deviations, both safe over-estimates or stricter: +//! * `stop_str_max_len` is the stop string's **UTF-8 byte length** — a provably +//! safe over-estimate of its token length (a token spans ≥ 1 byte, so +//! `bytes ≥ tokens`; `chars` is *not* a bound — one char can be several +//! tokens, e.g. `𓀀` → 3). The scheduler uses it only as a match-window +//! *size* (capped at the output length), so an over-estimate matches the same +//! stops — only an under-estimate misses. Python encodes each stop with the +//! tokenizer for the exact count; the byte bound avoids needing it here. +//! * `n > 1` (parallel sampling) is rejected — the rust egress maps one rid to +//! one response, so every sample past the first would be dropped. + +use serde::{Deserialize, Serialize}; + +use crate::error::Error; + +/// The sampling parameters of one `/generate` request. Deserialized from the +/// client's `sampling_params` object (unknown keys are a 400, mirroring Python's +/// `SamplingParams(**kwargs)` TypeError) and serialized by field name into the +/// scheduler header once [`normalize`](Self::normalize) has run. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[serde(deny_unknown_fields)] +pub struct SamplingParams { + pub max_new_tokens: Option, + pub temperature: f64, +} + +impl SamplingParams { + /// Normalize the fields, applying defaults and coercing types. Mirrors + /// Python's `SamplingParams.__post_init__` → `normalize`. + pub fn normalize( + &mut self, + _skip_tokenizer_init: bool, + _vocab_size: Option, + ) -> Result<(), Error> { + todo!() + } + + /// Verify the normalized fields are in range. Mirrors Python's + /// `SamplingParams.verify`. + pub fn verify(&self) -> Result<(), String> { + todo!() + } + + pub fn max_tokens_len(&self) -> usize { + todo!() + } +} + +/// The `/generate` body's `sampling_params`: one object (broadcast to every +/// prompt) or a list of them (one per prompt), fanned out by `GenerateBody::into_requests`. +/// +/// Hand-written `Deserialize` rather than `#[serde(untagged)]`: untagged buffers +/// the input and, on failure, reports only "data did not match any variant" — +/// losing the field-level message ("unknown field `temperature`, expected one of +/// …") that makes a typo actionable. Object-vs-list is unambiguous here, so a +/// single `deserialize_any` dispatch keeps the inner error verbatim. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum SamplingParamsInput { + /// Boxed: `SamplingParams` is ~440 bytes, so an inline variant would make + /// every `GenerateBody` that big regardless of which form arrived. + One(Box), + Many(Vec), +} diff --git a/rust/sglang-server/src/message/types.rs b/rust/sglang-server/src/message/types.rs new file mode 100644 index 000000000..c15027555 --- /dev/null +++ b/rust/sglang-server/src/message/types.rs @@ -0,0 +1,265 @@ +//! Shared wire-shape types: the token-id buffer alias, the scalar-or-list +//! adapter (with the sealed allowlist that keeps its `untagged` selection safe), +//! and the msgspec tagging machinery behind [`wire_struct!`]. + +use serde::{Deserialize, Serialize}; + +/// A flat token-id buffer — one request's `array("q")` cell on the Python side. +/// Wrapped in [`OneOrMany`] on the wire, where a bare list is one prompt's ids +/// (or a broadcast) and a list of lists is per-prompt. +pub type TokenIds = Vec; + +/// A field taking a bare `T` **or** `[T,…]` (`text: "hi"` or `text: ["a","b"]`). +/// `untagged` takes the first variant that matches, so a `T` that itself accepts +/// a sequence would make `Many` unreachable — hence the [`OneOrManyItem`] gate. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum OneOrMany { + One(T), + Many(Vec), +} + +/// Types vetted for [`OneOrMany`]. Sealed, so adding one is a deliberate act in +/// this file. +/// +/// **Never implement this for a self-describing type — `serde_json::Value`, +/// `rmpv::Value`, or anything else that deserializes from *any* shape.** Such a +/// `T` matches `[1,2]` as `One(Value::Array(…))`, so `Many` is never selected and +/// a batch silently arrives as a single request. Those types need a +/// `deserialize_any` dispatch instead (see [`SamplingParamsInput`]). +/// +/// [`TokenIds`] is the one member that does accept a sequence, and that ambiguity +/// is the intended semantics: flat `[1,2]` is one prompt's ids (or a broadcast), +/// `[[1],[2]]` is per-prompt — the shapes Python's `_normalize_batch` +/// distinguishes. `String` / `bool` / `i64` never match a list, so both forms +/// round-trip. +pub trait OneOrManyItem: sealed::SealedItem {} + +impl OneOrManyItem for T {} + +mod sealed { + /// Supertrait no downstream module can implement, sealing [`super::OneOrManyItem`]. + pub trait SealedItem {} + + impl SealedItem for bool {} + impl SealedItem for i64 {} + impl SealedItem for String {} + impl SealedItem for super::TokenIds {} +} + +/// A msgspec `tag=True` struct: element 0 of its array is the Python class name +/// the scheduler decodes by. Declared explicitly rather than taken from +/// `type_name`, which is unspecified and path-qualified. +pub(super) trait Tagged { + const TAG: &'static str; +} + +/// Declare msgspec `array_like=True` wire structs by their *own* fields, in wire +/// order; the inherited `BaseReq` preamble (`tag`, `rid`, `http_worker_ipc`) and +/// the [`Tagged`] impl are generated. **The struct name is the Python class +/// name** (via `stringify!`), so name it exactly as `io_struct.py` does — the +/// per-message tests assert the tag, so a rename fails loudly. Field attributes +/// pass through. The http_worker_ipc is just a placeholder for the scheduler's +/// `BaseReq` and is always `()`, without it the scheduler's `BaseReq` would be +/// misaligned on the wire. +/// +/// Two forms, and **one invocation may use only one** (an arm matches the whole +/// invocation): `Name<'a>` borrows its rid — zero-copy, for the hot path; +/// `Name` owns it, for a message held by the owned `Request` it would borrow. +macro_rules! wire_struct { + ($( + $(#[$meta:meta])* + $vis:vis $name:ident<$lt:lifetime> { + $($(#[$field_meta:meta])* $field:ident: $ty:ty,)* + } + )+) => {$( + $(#[$meta])* + #[derive(Debug)] + $vis struct $name<$lt> { + rid: &$lt str, + $($(#[$field_meta])* $field: $ty,)* + } + + /// Hand-written so the `BaseReq` preamble is SYNTHESIZED rather than stored. + /// + /// `tag` and `http_worker_ipc` are the same two values for every message, so + /// carrying them as fields meant every constructor restated them — and a + /// `Default`-based shortcut is the wrong fix twice over: the borrowed form + /// holds `&SamplingParams`, which has no `Default` at all, and + /// `..Default::default()` would let a field added here but missed in a + /// constructor ship a silent default on a POSITIONAL wire. Emitting them + /// here instead means the tag comes from [`Tagged::TAG`] and cannot be + /// forgotten, mistyped, or paired with the wrong struct. + /// + /// `serialize_struct` (not `serialize_seq`) keeps `rmp_serde` on exactly the + /// code path the derive used, so the bytes are unchanged — which + /// `to_header_msgpack_is_positionally_aligned` asserts index by index. + impl<$lt> Serialize for $name<$lt> { + fn serialize(&self, serializer: S) -> Result { + use serde::ser::SerializeStruct; // codespell:ignore ser + // Annotated so the zero-field case (a bare `BaseReq`) still infers. + let own = <[&'static str]>::len(&[$(stringify!($field)),*]); + let mut st = serializer.serialize_struct(stringify!($name), 3 + own)?; + st.serialize_field("tag", ::TAG)?; + st.serialize_field("rid", &self.rid)?; + st.serialize_field("http_worker_ipc", &())?; + $(st.serialize_field(stringify!($field), &self.$field)?;)* + st.end() + } + } + + impl<$lt> $name<$lt> { + pub fn encode(&self) -> Result { + rmp_serde::to_vec(self) + .map(Bytes::from) + .map_err(|e| Error::Codec(e.to_string())) + } + } + + impl<$lt> Tagged for $name<$lt> { + const TAG: &'static str = stringify!($name); + } + )+}; + ($( + $(#[$meta:meta])* + $vis:vis $name:ident { + $($(#[$field_meta:meta])* $field:ident: $ty:ty,)* + } + )+) => {$( + $(#[$meta])* + #[derive(Debug)] + $vis struct $name { + rid: String, + $($(#[$field_meta])* $field: $ty,)* + } + + /// Hand-written so the `BaseReq` preamble is SYNTHESIZED rather than stored. + /// + /// `tag` and `http_worker_ipc` are the same two values for every message, so + /// carrying them as fields meant every constructor restated them — and a + /// `Default`-based shortcut is the wrong fix twice over: the borrowed form + /// holds `&SamplingParams`, which has no `Default` at all, and + /// `..Default::default()` would let a field added here but missed in a + /// constructor ship a silent default on a POSITIONAL wire. Emitting them + /// here instead means the tag comes from [`Tagged::TAG`] and cannot be + /// forgotten, mistyped, or paired with the wrong struct. + /// + /// `serialize_struct` (not `serialize_seq`) keeps `rmp_serde` on exactly the + /// code path the derive used, so the bytes are unchanged — which + /// `to_header_msgpack_is_positionally_aligned` asserts index by index. + impl Serialize for $name { + fn serialize(&self, serializer: S) -> Result { + use serde::ser::SerializeStruct; // codespell:ignore ser + // Annotated so the zero-field case (a bare `BaseReq`) still infers. + let own = <[&'static str]>::len(&[$(stringify!($field)),*]); + let mut st = serializer.serialize_struct(stringify!($name), 3 + own)?; + st.serialize_field("tag", ::TAG)?; + st.serialize_field("rid", &self.rid)?; + st.serialize_field("http_worker_ipc", &())?; + $(st.serialize_field(stringify!($field), &self.$field)?;)* + st.end() + } + } + + impl $name { + pub fn get_rid(&self) -> &str { + &self.rid + } + + pub fn encode(&self) -> Result { + rmp_serde::to_vec(self) + .map(Bytes::from) + .map_err(|e| Error::Codec(e.to_string())) + } + } + + impl Tagged for $name { + const TAG: &'static str = stringify!($name); + } + )+}; +} + +/// Declare every owned-rid control message *and* the `ControlRequest` enum that +/// carries them. The enum, its variants and its delegating methods are generated +/// from this one list, so adding a message is a single declaration instead of an +/// edit in three places (enum + each method). +macro_rules! control_messages { + ($( + $(#[$meta:meta])* + $name:ident { + $($(#[$field_meta:meta])* $field:ident: $ty:ty,)* + } + )+) => { + wire_struct! {$( + $(#[$meta])* + pub(crate) $name { + $($(#[$field_meta])* $field: $ty,)* + } + )+} + + /// Which control message a request carries, as the wire struct itself. + /// These own their rid (`String`), so the enum needs no lifetime — a + /// borrowed rid would point back at the `Request` that owns this value. + #[derive(Debug)] + pub enum ControlRequest { + $($name($name),)+ + } + + impl ControlRequest { + /// The rid this message carries; `submit` reuses it as the request's + /// rid, so the two cannot disagree. + pub(crate) fn rid(&self) -> &str { + match self { + $(Self::$name(m) => m.get_rid(),)+ + } + } + + /// Encode as the msgspec tagged array. The variant selects the wire + /// struct, so an unknown tag is not representable. + pub(crate) fn encode(&self) -> Result { + match self { + $(Self::$name(m) => m.encode(),)+ + } + } + } + }; +} + +// `macro_rules!` is scoped by declaration order; name it so siblings can `use` it. +pub(super) use {control_messages, wire_struct}; + +#[cfg(test)] +mod tests { + use super::*; + + /// Pins `untagged`'s first-match-wins variant selection for the vetted + /// [`OneOrManyItem`] types: the `TokenIds` rows are the shapes + /// `GenerateBody::into_requests` relies on (flat = one prompt / broadcast, nested = + /// per-prompt), and `String` is the unambiguous case. + #[test] + fn untagged_selects_the_first_matching_variant() { + let one_of = |json: &str| -> bool { + matches!( + serde_json::from_str::>(json).unwrap(), + OneOrMany::One(_) + ) + }; + assert!(one_of("[1,2]"), "a flat id list is one prompt's ids"); + assert!(!one_of("[[1],[2]]"), "a nested list is per-prompt"); + + // A string can never match a sequence, so both forms stay unambiguous. + assert!(matches!( + serde_json::from_str::>(r#""hi""#).unwrap(), + OneOrMany::One(_) + )); + assert!(matches!( + serde_json::from_str::>(r#"["a","b"]"#).unwrap(), + OneOrMany::Many(v) if v.len() == 2 + )); + + // The hazard case is no longer expressible: `OneOrMany` + // fails to compile because `Value` is not an `OneOrManyItem`, so `Many` + // can never be silently unreachable. (Verified by construction — adding + // that instantiation anywhere is a compile error.) + } +}