From b83a59835dcaabd0e606194df62a8b8601710cf9 Mon Sep 17 00:00:00 2001 From: Rain Jiang <96632942+rainj-me@users.noreply.github.com> Date: Tue, 8 Sep 2026 00:41:31 -0700 Subject: [PATCH] sglang-server remove opaque type (#38095) --- rust/sglang-server/src/api_server/prefetch.rs | 71 +-- rust/sglang-server/src/message.rs | 1 + rust/sglang-server/src/message/multimodal.rs | 424 ++++++++++++++++++ rust/sglang-server/src/message/request.rs | 254 ++++------- rust/sglang-server/src/message/types.rs | 13 +- .../src/multi_modality/payload.rs | 191 +++----- .../tokenizer_manager/to_scheduler_tests.rs | 10 +- 7 files changed, 638 insertions(+), 326 deletions(-) create mode 100644 rust/sglang-server/src/message/multimodal.rs diff --git a/rust/sglang-server/src/api_server/prefetch.rs b/rust/sglang-server/src/api_server/prefetch.rs index b011a9538..6a6a0dc95 100644 --- a/rust/sglang-server/src/api_server/prefetch.rs +++ b/rust/sglang-server/src/api_server/prefetch.rs @@ -18,7 +18,7 @@ use sglang_mm::driver::{MAX_ITEMS_PER_REQUEST, MAX_REQUEST_BYTES}; use tokio::sync::Semaphore; use crate::message::request::{GenerateRequest, MmData}; -use crate::multi_modality::payload::{io_sources, item_count}; +use crate::multi_modality::payload::io_sources; /// Global bound on concurrent media fetches across all in-flight requests; /// excess acquisitions queue on the semaphore without holding a thread. @@ -40,23 +40,22 @@ pub async fn prefetch_all( let Some(mm) = mm.as_deref() else { return Ok(Vec::new()); }; - let values = [ - ("image", mm.image_data.as_ref()), - ("video", mm.video_data.as_ref()), - ("audio", mm.audio_data.as_ref()), + let modalities = [ + ("image", &mm.image_data), + ("video", &mm.video_data), + ("audio", &mm.audio_data), ]; - let items = values + let items = modalities .iter() - .filter_map(|(_, value)| *value) - .map(item_count) + .map(|(_, items)| items.len()) .sum::(); if items > MAX_ITEMS_PER_REQUEST { return Err(format!( "multimodal request exceeds {MAX_ITEMS_PER_REQUEST} media items" )); } - for (modality, value) in values { - let count = value.map(item_count).unwrap_or_default(); + for (modality, items) in modalities { + let count = items.len(); if let Some(limit) = modality_limits.get(modality) && count > *limit { @@ -66,10 +65,9 @@ pub async fn prefetch_all( )); } } - Ok(values + Ok(modalities .iter() - .filter_map(|(_, value)| *value) - .flat_map(io_sources) + .flat_map(|(_, items)| io_sources(items)) .collect()) }; let plans = requests @@ -120,9 +118,12 @@ async fn fetch_ordered(sources: Vec, total_bytes: u64) -> Result) -> MmItem { + MmItem::Source(s.into()) + } fn serve(bodies: Vec>) -> std::net::SocketAddr { let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); @@ -149,10 +150,10 @@ mod tests { addr } - fn mm_request(image_data: Value) -> GenerateRequest { + fn mm_request(image_data: Vec) -> GenerateRequest { GenerateRequest { mm: Some(Box::new(MmData { - image_data: Some(image_data), + image_data, ..Default::default() })), ..Default::default() @@ -173,9 +174,9 @@ mod tests { } let mut requests = vec![GenerateRequest { mm: Some(Box::new(MmData { - image_data: Some(Value::from(paths[0].display().to_string())), - video_data: Some(Value::from(paths[1].display().to_string())), - audio_data: Some(Value::from(paths[2].display().to_string())), + image_data: vec![src(paths[0].display().to_string())], + video_data: vec![src(paths[1].display().to_string())], + audio_data: vec![src(paths[2].display().to_string())], ..Default::default() })), ..Default::default() @@ -197,12 +198,12 @@ mod tests { let path = std::env::temp_dir().join(format!("sglang-prefetch-{}", std::process::id())); std::fs::write(&path, b"zzz").unwrap(); let mut requests = vec![ - mm_request(Value::Array(vec![ - Value::from(format!("http://{addr}/a.png")), - Value::from("data:image/png;base64,x"), - Value::from(format!("http://{addr}/b.png")), - Value::from(path.display().to_string()), - ])), + mm_request(vec![ + src(format!("http://{addr}/a.png")), + src("data:image/png;base64,x"), + src(format!("http://{addr}/b.png")), + src(path.display().to_string()), + ]), GenerateRequest::default(), ]; prefetch_all(&mut requests, &BTreeMap::new()).await.unwrap(); @@ -218,7 +219,7 @@ mod tests { #[tokio::test] async fn failed_download_rejects() { - let mut requests = vec![mm_request(Value::from("http://127.0.0.1:1/nope.png"))]; + let mut requests = vec![mm_request(vec![src("http://127.0.0.1:1/nope.png")])]; let err = prefetch_all(&mut requests, &BTreeMap::new()) .await .err() @@ -230,10 +231,10 @@ mod tests { /// fail to fetch, so a fetch error would prove fetching started. #[tokio::test] async fn item_budget_rejects_before_fetching() { - let sources: Vec = (0..=MAX_ITEMS_PER_REQUEST) - .map(|i| Value::from(format!("/definitely/not/here-{i}.png"))) + let sources: Vec = (0..=MAX_ITEMS_PER_REQUEST) + .map(|i| src(format!("/definitely/not/here-{i}.png"))) .collect(); - let mut requests = vec![mm_request(Value::Array(sources))]; + let mut requests = vec![mm_request(sources)]; let err = prefetch_all(&mut requests, &BTreeMap::new()) .await .err() @@ -249,11 +250,11 @@ mod tests { async fn per_modality_budget_rejects_before_fetching() { let mut requests = vec![GenerateRequest { mm: Some(Box::new(MmData { - image_data: Some(Value::Array(vec![ - Value::from("/definitely/not/here-0.png"), - Value::from("/definitely/not/here-1.png"), - ])), - video_data: Some(Value::Array(vec![Value::from("/definitely/not/here.mp4")])), + image_data: vec![ + src("/definitely/not/here-0.png"), + src("/definitely/not/here-1.png"), + ], + video_data: vec![src("/definitely/not/here.mp4")], ..Default::default() })), ..Default::default() diff --git a/rust/sglang-server/src/message.rs b/rust/sglang-server/src/message.rs index ff5904462..b35df6fc7 100644 --- a/rust/sglang-server/src/message.rs +++ b/rust/sglang-server/src/message.rs @@ -7,6 +7,7 @@ pub mod detok; pub mod finish_reason; pub mod ids; pub mod io_struct; +pub mod multimodal; pub mod request; pub mod response; pub mod sampling; diff --git a/rust/sglang-server/src/message/multimodal.rs b/rust/sglang-server/src/message/multimodal.rs new file mode 100644 index 000000000..e7da19d46 --- /dev/null +++ b/rust/sglang-server/src/message/multimodal.rs @@ -0,0 +1,424 @@ +//! Typed multimodal inputs of the `/generate` body — the Rust form of Python +//! `MultimodalDataInputFormat` (`io_struct.py`) — and their per-request fan-out. + +use std::fmt; + +use serde::de::value::{MapAccessDeserializer, SeqAccessDeserializer}; +use serde::de::{MapAccess, SeqAccess, Visitor}; +use serde::{Deserialize, Deserializer}; + +use super::request::{HeapBytes, check_broadcast_budget}; +use crate::utils::error::Error; + +/// One media item: Python `MultimodalDataInputItem` as it can arrive over JSON. +/// `bytes` and PIL images exist only on the in-process Engine path, so they have +/// no variant here. +#[derive(Debug, Clone, PartialEq)] +pub enum MmItem { + /// URL, `file://` / absolute path, `data:` URI, or bare base64 (Python `str`). + Source(String), + /// Python `ImageData` / `VideoData` (`{"url": …, …}`). Only `url` is kept: + /// the hint keys (`detail`, `max_dynamic_patch`, `preprocess_kwargs`, ...) + /// are read by model families this pipeline does not run, and Python's + /// `load_image` itself reduces the item to `.url`. + Ref { url: String }, + /// A preprocessed item (`{"format": "processor_output" | "precomputed_embedding", …}`). + /// Parsed only far enough to be rejected by name at the MM stage; Python + /// ignores it the same way on a text-only model. + Preprocessed { format: String }, +} + +impl MmItem { + /// The raw source string for the modality pipeline, `None` for a + /// preprocessed item. + pub fn source(&self) -> Option<&str> { + match self { + MmItem::Source(source) | MmItem::Ref { url: source } => Some(source), + MmItem::Preprocessed { .. } => None, + } + } +} + +impl HeapBytes for MmItem { + fn heap_bytes(&self) -> usize { + match self { + MmItem::Source(s) | MmItem::Ref { url: s } | MmItem::Preprocessed { format: s } => { + s.len() + } + } + } +} + +/// The object form of an item, as Python's `Dict[str, Any]`: `format` marks a +/// preprocessed item (checked first, as `glm4v` does), `url` an `ImageData`. +#[derive(Deserialize)] +struct ItemObject { + #[serde(default)] + url: Option, + #[serde(default)] + format: Option, +} + +impl TryFrom for MmItem { + type Error = &'static str; + + fn try_from(object: ItemObject) -> Result { + match (object.format, object.url) { + (Some(format), _) => Ok(MmItem::Preprocessed { format }), + (None, Some(url)) => Ok(MmItem::Ref { url }), + (None, None) => Err("a multimodal item object needs a `url` or a `format` key"), + } + } +} + +/// Hand-written rather than `#[serde(untagged)]` so a bad item is reported as +/// what it is ("expected a source string or an item object"), not as "did not +/// match any variant". +impl<'de> Deserialize<'de> for MmItem { + fn deserialize>(deserializer: D) -> Result { + struct ItemVisitor; + + impl<'de> Visitor<'de> for ItemVisitor { + type Value = MmItem; + + fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("a media source string or a multimodal item object") + } + + fn visit_str(self, value: &str) -> Result { + Ok(MmItem::Source(value.to_owned())) + } + + fn visit_string(self, value: String) -> Result { + Ok(MmItem::Source(value)) + } + + fn visit_map>(self, map: A) -> Result { + ItemObject::deserialize(MapAccessDeserializer::new(map))? + .try_into() + .map_err(serde::de::Error::custom) + } + } + + deserializer.deserialize_any(ItemVisitor) + } +} + +/// One `image_data` / `video_data` / `audio_data` field as sent: Python +/// `MultimodalDataInputFormat`, whose three shapes read differently for a single +/// request and a batch (see [`fan_out`]). +#[derive(Debug, Clone, PartialEq)] +pub enum MmDataInput { + /// One item: a single request's whole input, or a broadcast to every batch entry. + One(MmItem), + /// A flat list: a single request's items, or one item per batch entry. + Many(Vec>), + /// One item list per batch entry. + Nested(Vec>>>), +} + +/// One element of a list-form field, before the list is known to be flat or nested. +enum ListElement { + Null, + Item(MmItem), + List(Vec>), +} + +impl<'de> Deserialize<'de> for ListElement { + fn deserialize>(deserializer: D) -> Result { + struct ElementVisitor; + + impl<'de> Visitor<'de> for ElementVisitor { + type Value = ListElement; + + fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("null, a media source string, an item object, or a list of items") + } + + fn visit_unit(self) -> Result { + Ok(ListElement::Null) + } + + fn visit_none(self) -> Result { + Ok(ListElement::Null) + } + + fn visit_str(self, value: &str) -> Result { + Ok(ListElement::Item(MmItem::Source(value.to_owned()))) + } + + fn visit_string(self, value: String) -> Result { + Ok(ListElement::Item(MmItem::Source(value))) + } + + fn visit_map>(self, map: A) -> Result { + ItemObject::deserialize(MapAccessDeserializer::new(map))? + .try_into() + .map(ListElement::Item) + .map_err(serde::de::Error::custom) + } + + fn visit_seq>(self, seq: A) -> Result { + Vec::>::deserialize(SeqAccessDeserializer::new(seq)) + .map(ListElement::List) + } + } + + deserializer.deserialize_any(ElementVisitor) + } +} + +impl<'de> Deserialize<'de> for MmDataInput { + fn deserialize>(deserializer: D) -> Result { + struct InputVisitor; + + impl<'de> Visitor<'de> for InputVisitor { + type Value = MmDataInput; + + fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("a media item, a list of items, or a list of item lists") + } + + fn visit_str(self, value: &str) -> Result { + Ok(MmDataInput::One(MmItem::Source(value.to_owned()))) + } + + fn visit_string(self, value: String) -> Result { + Ok(MmDataInput::One(MmItem::Source(value))) + } + + fn visit_map>(self, map: A) -> Result { + ItemObject::deserialize(MapAccessDeserializer::new(map))? + .try_into() + .map(MmDataInput::One) + .map_err(serde::de::Error::custom) + } + + fn visit_seq>(self, seq: A) -> Result { + let elements = Vec::::deserialize(SeqAccessDeserializer::new(seq))?; + let nested = elements + .iter() + .any(|element| matches!(element, ListElement::List(_))); + if !nested { + return Ok(MmDataInput::Many( + elements + .into_iter() + .map(|element| match element { + ListElement::Item(item) => Some(item), + ListElement::Null => None, + ListElement::List(_) => unreachable!("checked above"), + }) + .collect(), + )); + } + elements + .into_iter() + .map(|element| match element { + ListElement::List(items) => Ok(Some(items)), + ListElement::Null => Ok(None), + ListElement::Item(_) => Err(serde::de::Error::custom( + "a nested list cannot mix bare items with item lists", + )), + }) + .collect::>() + .map(MmDataInput::Nested) + } + } + + deserializer.deserialize_any(InputVisitor) + } +} + +/// The items of one modality for one request, `null` entries dropped. +fn present(items: Vec>) -> Vec { + items.into_iter().flatten().collect() +} + +/// Fan one field into per-request item lists (empty = no input for that +/// request), mirroring Python `_normalize_{image,video,audio}_data`: +/// * absent, `[]`, or all-`null` → no input (Python `has_valid_data`); +/// * single request → one item or a flat list, taken as is; +/// * batch + one item → broadcast to every entry; +/// * batch + list → per entry, length must equal the batch size. +/// +/// The Python image path wraps a broadcast as `[[img]] * num` while video and +/// audio broadcast bare; the difference vanishes here because every request's +/// input is already an item list. +pub fn fan_out( + value: Option, + n: usize, + is_batch: bool, + name: &str, +) -> Result>, Error> { + let Some(value) = value else { + return Ok(vec![Vec::new(); n]); + }; + if !is_batch { + return match value { + MmDataInput::One(item) => Ok(vec![vec![item]]), + MmDataInput::Many(items) => Ok(vec![present(items)]), + MmDataInput::Nested(_) => Err(Error::Validation(format!( + "{name}: a nested list is the batch form; a single request takes one item or a flat list" + ))), + }; + } + match value { + MmDataInput::One(item) => { + // A broadcast deep-clones once per prompt — same blow-up as + // sampling_params, so bound the product before any clone. + check_broadcast_budget(item.heap_bytes(), n, name)?; + Ok(vec![vec![item]; n]) + } + MmDataInput::Many(items) if items.is_empty() => Ok(vec![Vec::new(); n]), + MmDataInput::Many(items) => { + check_len(items.len(), n, name)?; + Ok(items + .into_iter() + .map(|item| item.into_iter().collect()) + .collect()) + } + MmDataInput::Nested(lists) => { + check_len(lists.len(), n, name)?; + Ok(lists + .into_iter() + .map(|items| items.map(present).unwrap_or_default()) + .collect()) + } + } +} + +fn check_len(len: usize, n: usize, name: &str) -> Result<(), Error> { + if len != n { + return Err(Error::Validation(format!( + "{name}: list length {len} does not match batch size {n}" + ))); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn parse(json: &str) -> Result { + serde_json::from_str(json) + } + + fn src(s: &str) -> MmItem { + MmItem::Source(s.to_owned()) + } + + /// The three Python shapes parse to their own variants, with `null` + /// entries kept in place so batch fan-out can index them. + #[test] + fn parses_python_shapes() { + assert_eq!(parse(r#""u""#).unwrap(), MmDataInput::One(src("u"))); + assert_eq!( + parse(r#"["a", null, "b"]"#).unwrap(), + MmDataInput::Many(vec![Some(src("a")), None, Some(src("b"))]) + ); + assert_eq!(parse("[]").unwrap(), MmDataInput::Many(vec![])); + assert_eq!( + parse(r#"[["a", null], null, []]"#).unwrap(), + MmDataInput::Nested(vec![Some(vec![Some(src("a")), None]), None, Some(vec![])]) + ); + } + + /// Object items: `format` wins over `url` (a preprocessed item may carry + /// both), and an object with neither is named in the error. + #[test] + fn parses_item_objects() { + assert_eq!( + parse(r#"{"url": "u", "detail": "high"}"#).unwrap(), + MmDataInput::One(MmItem::Ref { url: "u".into() }) + ); + assert_eq!( + parse(r#"[{"format": "processor_output", "url": "u", "pixel_values": [1]}]"#).unwrap(), + MmDataInput::Many(vec![Some(MmItem::Preprocessed { + format: "processor_output".into() + })]) + ); + let err = parse(r#"{"detail": "high"}"#).unwrap_err().to_string(); + assert!(err.contains("`url` or a `format`"), "{err}"); + } + + /// Anything Python's item union does not cover is rejected up front, with + /// the expected shape in the message. + #[test] + fn rejects_non_items() { + for (json, expect) in [ + ("5", "expected a media item, a list of items"), + (r#"["a", 5]"#, "expected null, a media source string"), + (r#"["a", ["b"]]"#, "cannot mix"), + ( + r#"[[["a"]]]"#, + "expected a media source string or a multimodal item object", + ), + ] { + let err = parse(json).unwrap_err().to_string(); + assert!(err.contains(expect), "{json}: {err}"); + } + } + + #[test] + fn single_request_takes_item_or_flat_list() { + assert_eq!(fan_out(None, 1, false, "image_data").unwrap(), vec![vec![]]); + assert_eq!( + fan_out(Some(MmDataInput::One(src("u"))), 1, false, "image_data").unwrap(), + vec![vec![src("u")]] + ); + assert_eq!( + fan_out( + Some(parse(r#"["a", null, "b"]"#).unwrap()), + 1, + false, + "image_data" + ) + .unwrap(), + vec![vec![src("a"), src("b")]] + ); + assert_eq!( + fan_out(Some(parse("[null]").unwrap()), 1, false, "image_data").unwrap(), + vec![vec![]] + ); + let err = fan_out(Some(parse(r#"[["a"]]"#).unwrap()), 1, false, "image_data").unwrap_err(); + assert!(err.to_string().contains("batch form"), "{err}"); + } + + #[test] + fn batch_broadcasts_scalar_and_splits_lists() { + let one = fan_out(Some(MmDataInput::One(src("u"))), 2, true, "video_data").unwrap(); + assert_eq!(one, vec![vec![src("u")], vec![src("u")]]); + + let flat = fan_out( + Some(parse(r#"["a", null]"#).unwrap()), + 2, + true, + "image_data", + ) + .unwrap(); + assert_eq!(flat, vec![vec![src("a")], vec![]]); + + let nested = fan_out( + Some(parse(r#"[["a", "b"], null, [null]]"#).unwrap()), + 3, + true, + "image_data", + ) + .unwrap(); + assert_eq!(nested, vec![vec![src("a"), src("b")], vec![], vec![]]); + + // `[]` is "no input", not a length-0 per-entry list. + assert_eq!( + fan_out(Some(parse("[]").unwrap()), 2, true, "image_data").unwrap(), + vec![vec![], vec![]] + ); + for json in [r#"["a"]"#, r#"[["a"]]"#] { + let err = fan_out(Some(parse(json).unwrap()), 2, true, "image_data").unwrap_err(); + assert!( + err.to_string().contains("does not match batch size"), + "{json}: {err}" + ); + } + } +} diff --git a/rust/sglang-server/src/message/request.rs b/rust/sglang-server/src/message/request.rs index 10df74b94..8261103c3 100644 --- a/rust/sglang-server/src/message/request.rs +++ b/rust/sglang-server/src/message/request.rs @@ -9,6 +9,7 @@ use itertools::izip; use serde::Deserialize; use super::io_struct::{ControlRequest, TokenizedGenerateReqInput}; +use super::multimodal::{self, MmDataInput, MmItem}; use super::response::ResponseSink; use super::sampling::{SamplingParams, SamplingParamsInput}; use super::types::{OneOrMany, OneOrManyItem, TokenIds}; @@ -93,19 +94,17 @@ pub struct GenerateBody { /// DP routing hints — per-request scalars even for batches, as in Python. pub routed_dp_rank: Option, pub disagg_prefill_dp_rank: Option, - // Multimodal inputs, permissive `Value` so any shape Python's - // `GenerateReqInput` accepts (URL / base64 / list / list-of-lists) parses. - // `into_requests` fans them out per the Python - // `_normalize_{image,video,audio}_data` batch rules. - pub image_data: Option, + // Multimodal inputs (Python `MultimodalDataInputFormat`), fanned out per + // request by `multimodal::fan_out`. + pub image_data: Option, /// Caller-supplied per-item content hashes (hex) overriding the computed /// ones, so an external router's keys align with the prefix cache. Single - /// requests only: Python declares the batched shapes but `__getitem__` never - /// forwards them, so a batch is rejected here rather than answered with - /// hashes it did not ask for. - pub mm_hashes: Option, - pub video_data: Option, - pub audio_data: Option, + /// requests only: Python declares the batched (nested) shape but + /// `__getitem__` never forwards it, so a batch is rejected here rather than + /// answered with hashes it did not ask for. + pub mm_hashes: Option>>, + pub video_data: Option, + pub audio_data: Option, } impl GenerateBody { @@ -340,18 +339,26 @@ impl GenerateBody { // servers different prefix-cache keys for the same body. Reject instead of // dropping it silently as Python does — the field exists to align a // caller's keys, so ignoring it returns subtly wrong ones. - if is_batch && mm_value_present(&mm_hashes) { - return Err(Error::Validation( - "mm_hashes is not supported for batch requests; send one request per prompt".into(), - )); - } - // Multimodal columns; see `split_mm_column` for the Python parity rules. - let images = split_mm_column(image_data, n, is_batch, MmBroadcast::WrapInList) - .map_err(|e| Error::Validation(format!("image_data: {e}")))?; - let videos = split_mm_column(video_data, n, is_batch, MmBroadcast::AsIs) - .map_err(|e| Error::Validation(format!("video_data: {e}")))?; - let audios = split_mm_column(audio_data, n, is_batch, MmBroadcast::AsIs) - .map_err(|e| Error::Validation(format!("audio_data: {e}")))?; + let mm_hashes: Vec = match mm_hashes { + None => Vec::new(), + Some(OneOrMany::One(hashes)) if hashes.is_empty() => Vec::new(), + Some(_) if is_batch => { + return Err(Error::Validation( + "mm_hashes is not supported for batch requests; send one request per prompt" + .into(), + )); + } + Some(OneOrMany::One(hashes)) => hashes, + Some(OneOrMany::Many(_)) => { + return Err(Error::Validation( + "mm_hashes must be a flat list of hex strings for a single request".into(), + )); + } + }; + // Multimodal columns; see `multimodal::fan_out` for the Python parity rules. + let images = multimodal::fan_out(image_data, n, is_batch, "image_data")?; + let videos = multimodal::fan_out(video_data, n, is_batch, "video_data")?; + let audios = multimodal::fan_out(audio_data, n, is_batch, "audio_data")?; // 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. @@ -425,14 +432,8 @@ impl GenerateBody { .collect(); // Single requests only (batches rejected above). Malformed entries are // dropped here and warned about in `mm::apply_caller_hashes`, never a 400. - if !is_batch - && let (Some(rmpv::Value::Array(vals)), Some(req)) = (mm_hashes, requests.first_mut()) - && let Some(mm) = req.mm.as_deref_mut() - { - mm.mm_hashes = vals - .iter() - .filter_map(|v| v.as_str().map(str::to_owned)) - .collect(); + if let Some(mm) = requests.first_mut().and_then(|req| req.mm.as_deref_mut()) { + mm.mm_hashes = mm_hashes; } Ok((requests, is_batch)) } @@ -441,11 +442,11 @@ impl GenerateBody { /// Box the per-item mm values, `None` when the item has none — the common /// text-only case keeps `GenerateRequest` slim. fn pack_mm( - image_data: Option, - video_data: Option, - audio_data: Option, + image_data: Vec, + video_data: Vec, + audio_data: Vec, ) -> Option> { - if image_data.is_none() && video_data.is_none() && audio_data.is_none() { + if image_data.is_empty() && video_data.is_empty() && audio_data.is_empty() { return None; } Some(Box::new(MmData { @@ -456,60 +457,6 @@ fn pack_mm( })) } -/// How a scalar mm value broadcasts across a batch: images become a one-image -/// list per item (`[[img]] * num` in Python `_normalize_image_data`), -/// video/audio broadcast bare (`[v] * num` in `_normalize_video_data`). -#[derive(Clone, Copy)] -enum MmBroadcast { - WrapInList, - AsIs, -} - -/// Fan one mm field into per-item values, mirroring Python's -/// `_normalize_{image,video,audio}_data`: -/// * `None` / empty list → `None` for every item; -/// * single request → the raw value passes through (the processor wraps a -/// non-list into a one-element list); -/// * batch + scalar → broadcast to every item, per `MmBroadcast`; -/// * batch + list → per item, length must equal the batch size. -fn split_mm_column( - v: Option, - n: usize, - is_batch: bool, - broadcast: MmBroadcast, -) -> Result>, String> { - let Some(v) = v else { - return Ok(vec![None; n]); - }; - if v.is_nil() { - return Ok(vec![None; n]); - } - if !is_batch { - return Ok(vec![Some(v)]); - } - match v { - rmpv::Value::Array(items) if items.is_empty() => Ok(vec![None; n]), - rmpv::Value::Array(items) => { - if items.len() != n { - return Err(format!( - "list length {} does not match batch size {n}", - items.len() - )); - } - Ok(items.into_iter().map(Some).collect()) - } - scalar => { - // A broadcast deep-clones once per prompt — same blow-up as - // sampling_params above, so bound the product before any clone. - check_broadcast_budget(scalar.heap_bytes(), n, "value").map_err(|e| e.to_string())?; - Ok(match broadcast { - MmBroadcast::WrapInList => vec![Some(rmpv::Value::Array(vec![scalar])); n], - MmBroadcast::AsIs => vec![Some(scalar); n], - }) - } - } -} - /// One request handed to the MM worker pool: the rid to correlate the result, /// plus the owned inputs from [`GenerateRequest::take_mm_work`]. #[derive(Debug)] @@ -524,22 +471,15 @@ pub struct MmRequest { pub struct MmWorkItem { pub text: Option, pub input_ids: Option>, - pub image_data: Option, - pub video_data: Option, - pub audio_data: Option, + pub image_data: Vec, + pub video_data: Vec, + pub audio_data: Vec, /// See [`MmData::prefetched`]. pub prefetched: Vec, /// See [`GenerateBody::mm_hashes`]. pub mm_hashes: Vec, } -/// Whether an optional mm field counts as multimodal input, via the same -/// `value_present` the MM worker's payload parser uses. -fn mm_value_present(v: &Option) -> bool { - v.as_ref() - .is_some_and(crate::multi_modality::payload::value_present) -} - /// The owned request as it travels request stages (single owner, so `state` is /// mutated lock-free). Common fields here; variant data in [`RequestKind`]. #[derive(Debug)] @@ -649,22 +589,23 @@ pub struct GenerateRequest { /// so these are pure passthrough for the scheduler/LB protocol. pub routed_dp_rank: Option, pub disagg_prefill_dp_rank: Option, - /// Multimodal inputs, carried opaquely. Consumed by the Encoding stage, - /// which ships them to the MM worker pool; never read by the tokenizer or - /// serialized onto the scheduler header. Boxed so the common text-only - /// request doesn't grow every `Request` moved between stages. + /// Multimodal inputs. Consumed by the Encoding stage, which ships them to + /// the MM worker pool; never read by the tokenizer or serialized onto the + /// scheduler header. Boxed so the common text-only request doesn't grow + /// every `Request` moved between stages. pub mm: Option>, } -/// The opaque multimodal fields of one request (see [`GenerateRequest::mm`]). +/// The multimodal fields of one request (see [`GenerateRequest::mm`]), each +/// modality already fanned out to this request's own item list. /// /// Constructed directly only by tests: `api_server::prefetch` fills its /// `prefetched` field, everything else gets it packed inside a `GenerateRequest`. #[derive(Debug, Default)] pub struct MmData { - pub image_data: Option, - pub video_data: Option, - pub audio_data: Option, + pub image_data: Vec, + pub video_data: Vec, + pub audio_data: Vec, /// Bytes of `image_data`'s I/O-backed sources, resolved by /// `api_server::prefetch` in `payload::io_sources` order so MM workers /// never block on I/O. Out-of-band: the values above stay as the client @@ -684,9 +625,7 @@ impl GenerateRequest { /// Python `GenerateReqInput.contains_mm_input()`. pub fn has_multimodal(&self) -> bool { self.mm.as_ref().is_some_and(|mm| { - mm_value_present(&mm.image_data) - || mm_value_present(&mm.video_data) - || mm_value_present(&mm.audio_data) + !mm.image_data.is_empty() || !mm.video_data.is_empty() || !mm.audio_data.is_empty() }) } @@ -700,9 +639,9 @@ impl GenerateRequest { ..Default::default() }; if let Some(m) = self.mm.as_deref_mut() { - work.image_data = m.image_data.take(); - work.video_data = m.video_data.take(); - work.audio_data = m.audio_data.take(); + work.image_data = std::mem::take(&mut m.image_data); + work.video_data = std::mem::take(&mut m.video_data); + work.audio_data = std::mem::take(&mut m.audio_data); work.prefetched = std::mem::take(&mut m.prefetched); work.mm_hashes = std::mem::take(&mut m.mm_hashes); } @@ -730,7 +669,7 @@ impl GenerateRequest { /// 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 { +pub(super) trait HeapBytes { fn heap_bytes(&self) -> usize; } impl HeapBytes for bool { @@ -758,23 +697,6 @@ impl HeapBytes for Option { self.as_ref().map_or(0, HeapBytes::heap_bytes) } } -impl HeapBytes for rmpv::Value { - fn heap_bytes(&self) -> usize { - use rmpv::Value; - const NODE: usize = std::mem::size_of::(); - match self { - Value::String(s) => s.as_bytes().len(), - Value::Binary(b) => b.len(), - Value::Ext(_, b) => b.len(), - Value::Array(items) => items.iter().map(|v| NODE + v.heap_bytes()).sum(), - Value::Map(entries) => entries - .iter() - .map(|(k, v)| 2 * NODE + k.heap_bytes() + v.heap_bytes()) - .sum(), - _ => 0, - } - } -} /// Collapse `fan_out`'s nullable-element output: outer `None` (field absent / /// scalar broadcast of nothing) and inner `None` (an explicit `null` list @@ -784,7 +706,7 @@ fn flatten_column(column: Vec>>) -> Vec> { } /// Reject a broadcast whose clones would exceed [`MAX_BROADCAST_CLONE_BYTES`]. -fn check_broadcast_budget(per_clone: usize, n: usize, name: &str) -> Result<(), Error> { +pub(super) 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. @@ -999,32 +921,43 @@ mod tests { } /// Mm columns fan out per Python `_normalize_{image,video}_data`: a single - /// request passes the raw value through; a batch broadcasts a scalar image as - /// `[img]` per item, maps a list per item with matching lengths, and treats - /// `null`/`[]` as absent. + /// request keeps its items; a batch broadcasts a scalar to every item, maps + /// a list per item with matching lengths, and treats `null`/`[]` as absent. #[test] fn split_mm_fanout_matches_python_normalize() { - let image_of = |p: &GenerateRequest| p.mm.as_ref().unwrap().image_data.clone().unwrap(); + let src = |s: &str| MmItem::Source(s.to_owned()); + let images_of = |p: &GenerateRequest| p.mm.as_ref().unwrap().image_data.clone(); - // Single request: raw value passes through untouched. + // Single request: one item, or a flat list, kept as sent. let (ps, _) = requests(r#"{"text": "a", "image_data": "http://x/i.jpg"}"#).unwrap(); - assert_eq!(image_of(&ps[0]).as_str(), Some("http://x/i.jpg")); + assert_eq!(images_of(&ps[0]), vec![src("http://x/i.jpg")]); assert!(ps[0].has_multimodal()); + let (ps, _) = requests(r#"{"text": "a", "image_data": ["u1", {"url": "u2"}]}"#).unwrap(); + assert_eq!( + images_of(&ps[0]), + vec![src("u1"), MmItem::Ref { url: "u2".into() }] + ); - // Batch + scalar image: broadcast, wrapped as a one-image list per item. + // Batch + scalar image: broadcast, one image per item. let (ps, _) = requests(r#"{"text": ["a", "b"], "image_data": "u"}"#).unwrap(); for p in &ps { - assert_eq!(image_of(p).as_array().unwrap().len(), 1); + assert_eq!(images_of(p), vec![src("u")]); assert!(p.has_multimodal()); } - // Batch + per-item list: element i goes to item i. + // Batch + per-item list: element i goes to item i; nested lists are + // per-item lists. let (ps, _) = requests(r#"{"text": ["a", "b"], "image_data": ["u1", "u2"]}"#).unwrap(); - assert_eq!(image_of(&ps[0]).as_str(), Some("u1")); - assert_eq!(image_of(&ps[1]).as_str(), Some("u2")); + assert_eq!(images_of(&ps[0]), vec![src("u1")]); + assert_eq!(images_of(&ps[1]), vec![src("u2")]); + let (ps, _) = + requests(r#"{"text": ["a", "b"], "image_data": [["u1", "u2"], null]}"#).unwrap(); + assert_eq!(images_of(&ps[0]), vec![src("u1"), src("u2")]); + assert!(!ps[1].has_multimodal()); - // Batch + wrong-length list is a 400. + // Batch + wrong-length list is a 400, as is the batch shape on a single. assert!(requests(r#"{"text": ["a", "b"], "image_data": ["u1"]}"#).is_err()); + assert!(requests(r#"{"text": "a", "image_data": [["u1"]]}"#).is_err()); // null / [] mean "no multimodal input". let (ps, _) = requests(r#"{"text": "a", "image_data": null}"#).unwrap(); @@ -1032,11 +965,10 @@ mod tests { let (ps, _) = requests(r#"{"text": "a", "image_data": []}"#).unwrap(); assert!(!ps[0].has_multimodal()); - // Batch + scalar video: broadcast bare (not wrapped), per Python - // `_normalize_video_data`. + // Batch + scalar video broadcasts too (Python leaves it unwrapped, but + // every request's input is an item list here). let (ps, _) = requests(r#"{"text": ["a", "b"], "video_data": "v"}"#).unwrap(); - let video = ps[1].mm.as_ref().unwrap().video_data.clone().unwrap(); - assert_eq!(video.as_str(), Some("v")); + assert_eq!(ps[1].mm.as_ref().unwrap().video_data, vec![src("v")]); assert!(ps[1].has_multimodal()); } @@ -1045,17 +977,17 @@ mod tests { /// and are never charged. #[test] fn oversized_mm_broadcast_rejected() { - let big = rmpv::Value::from("x".repeat(MAX_BROADCAST_CLONE_BYTES / 2 + 1)); - let err = split_mm_column(Some(big.clone()), 2, true, MmBroadcast::WrapInList) + let big = MmItem::Source("x".repeat(MAX_BROADCAST_CLONE_BYTES / 2 + 1)); + let err = multimodal::fan_out(Some(MmDataInput::One(big.clone())), 2, true, "image_data") .err() .unwrap(); - assert!(err.contains("broadcast"), "{err}"); + assert!(err.to_string().contains("broadcast"), "{err}"); // A per-item list of the same total size moves, not clones: accepted. - let list = rmpv::Value::Array(vec![big, rmpv::Value::from("y")]); - assert!(split_mm_column(Some(list), 2, true, MmBroadcast::WrapInList).is_ok()); + let list = MmDataInput::Many(vec![Some(big), Some(MmItem::Source("y".into()))]); + assert!(multimodal::fan_out(Some(list), 2, true, "image_data").is_ok()); // Small scalars broadcast fine. - let small = rmpv::Value::from("u1"); - assert!(split_mm_column(Some(small), 2, true, MmBroadcast::AsIs).is_ok()); + let small = MmDataInput::One(MmItem::Source("u1".into())); + assert!(multimodal::fan_out(Some(small), 2, true, "audio_data").is_ok()); } /// `mm_hashes` rides only on single requests (Python `__getitem__` @@ -1068,10 +1000,12 @@ mod tests { assert_eq!(ps[0].take_mm_work().mm_hashes, vec!["a1b2", "0xff"]); assert!(ps[0].mm.as_ref().unwrap().mm_hashes.is_empty()); - // A batch cannot carry hashes (Python drops them), so it is rejected... + // A batch cannot carry hashes (Python drops them), so it is rejected, + // as is the nested batch shape on a single request... for body in [ r#"{"text": ["a", "b"], "image_data": ["u", "v"], "mm_hashes": [["x"], ["y"]]}"#, r#"{"text": ["a", "b"], "image_data": ["u", "v"], "mm_hashes": ["x", "y"]}"#, + r#"{"text": "a", "image_data": "u", "mm_hashes": [["x"]]}"#, ] { let err = requests(body).err().unwrap(); assert!(matches!(err, Error::Validation(_)), "{body}: {err:?}"); @@ -1094,11 +1028,11 @@ mod tests { let work = ps[0].take_mm_work(); assert_eq!(work.text.as_deref(), Some("hi")); assert!(work.input_ids.is_none()); - assert_eq!(work.image_data.unwrap().as_array().unwrap().len(), 2); - assert!(work.video_data.is_none()); - assert_eq!(work.audio_data.unwrap().as_str(), Some("a")); + assert_eq!(work.image_data.len(), 2); + assert!(work.video_data.is_empty()); + assert_eq!(work.audio_data, vec![MmItem::Source("a".into())]); // Moved out, not cloned; `text` survives for the header. - assert!(ps[0].mm.as_ref().unwrap().image_data.is_none()); + assert!(ps[0].mm.as_ref().unwrap().image_data.is_empty()); assert_eq!(ps[0].text.as_deref(), Some("hi")); } diff --git a/rust/sglang-server/src/message/types.rs b/rust/sglang-server/src/message/types.rs index 360d8f5af..e01fae2a8 100644 --- a/rust/sglang-server/src/message/types.rs +++ b/rust/sglang-server/src/message/types.rs @@ -28,11 +28,12 @@ pub enum OneOrMany { /// 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. +/// [`TokenIds`] and `Vec` are the members that do 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, and `mm_hashes`'s +/// `Union[List[str], List[List[str]]]` reads the same way. `String` / `bool` / +/// `i64` never match a list, so both forms round-trip. pub trait OneOrManyItem: sealed::SealedItem {} impl OneOrManyItem for T {} @@ -45,6 +46,8 @@ mod sealed { impl SealedItem for i64 {} impl SealedItem for String {} impl SealedItem for super::TokenIds {} + /// `mm_hashes`: a flat list is one request's hashes, nested is per-request. + impl SealedItem for Vec {} // Nullable elements for the PD bootstrap fields (`List[Optional[...]]` in // Python — the PD router sends `bootstrap_port: [null, …]` when deferring to // the scheduler's default port). A bare `null` never reaches `One(None)`: the diff --git a/rust/sglang-server/src/multi_modality/payload.rs b/rust/sglang-server/src/multi_modality/payload.rs index f047c3bf7..7c2d3277f 100644 --- a/rust/sglang-server/src/multi_modality/payload.rs +++ b/rust/sglang-server/src/multi_modality/payload.rs @@ -6,16 +6,16 @@ //! precomputed features, …). use bytes::Bytes; -use rmpv::Value; use sglang_mm::driver::{ImageSource, MmInput}; +use crate::message::multimodal::MmItem; use crate::message::request::MmWorkItem; /// True for sources the API layer must resolve before MM dispatch: I/O — network /// *or* disk, since a network mount can hang past any HTTP timeout — never runs /// on the fixed MM worker pool (see `api_server::prefetch`). `data:` and bare -/// base64 are pure CPU and stay on the worker. Lives next to `collect_images` so -/// the prefetch walk and the parse walk cannot drift. +/// base64 are pure CPU and stay on the worker. Lives next to [`image_source`] +/// so the prefetch walk and the parse walk cannot drift. pub fn is_io_source(src: &str) -> bool { src.starts_with("http://") || src.starts_with("https://") @@ -23,103 +23,64 @@ pub fn is_io_source(src: &str) -> bool { || src.starts_with('/') } -/// The I/O-backed sources of an `image_data` value, in `collect_images` order. -pub fn io_sources(value: &Value) -> Vec { - let mut out = Vec::new(); - let mut walk = |value: &Value| { - if let Some(src) = value.as_str().filter(|s| is_io_source(s)) { - out.push(src.to_owned()); - } - }; - if let Value::Array(values) = value { - values.iter().for_each(&mut walk); - } else { - walk(value); - } - out -} - -/// How many media items an `image_data` value contributes, walked the way -/// [`collect_images`] walks it, so the item budget can reject before fetching. -pub fn item_count(value: &Value) -> usize { - match value { - Value::Nil => 0, - Value::Array(values) => values.iter().map(item_count).sum(), - _ => 1, - } +/// The I/O-backed sources of one modality's items, in item order. +pub fn io_sources(items: &[MmItem]) -> Vec { + items + .iter() + .filter_map(MmItem::source) + .filter(|src| is_io_source(src)) + .map(str::to_owned) + .collect() } /// I/O-backed sources are swapped for their `work.prefetched` bytes (in /// [`io_sources`] order); one left without an entry is an internal error here, /// never a fetch. pub fn to_mm_input(work: MmWorkItem) -> Result { - let present = |v: &Option| v.as_ref().is_some_and(value_present); - if present(&work.video_data) || present(&work.audio_data) { + let MmWorkItem { + text, + input_ids, + image_data, + video_data, + audio_data, + prefetched, + mm_hashes: _, + } = work; + if !video_data.is_empty() || !audio_data.is_empty() { return Err("unsupported modality: video/audio input".into()); } - let mut images = Vec::new(); - if let Some(image_data) = &work.image_data { - collect_images(image_data, &mut work.prefetched.iter(), &mut images)?; - } + let mut prefetched = prefetched.iter(); + let images = image_data + .into_iter() + .map(|item| image_source(item, &mut prefetched)) + .collect::, _>>()?; if images.is_empty() { return Err("no raw image sources in mm input".into()); } Ok(MmInput { - text: work.text, - input_ids: work.input_ids, + text, + input_ids, images, }) } -fn collect_images( - value: &Value, +fn image_source( + item: MmItem, prefetched: &mut std::slice::Iter, - out: &mut Vec, -) -> Result<(), String> { - match value { - Value::Nil => Ok(()), - Value::String(value) => { - let value = value - .as_str() - .ok_or_else(|| "non-utf8 image source".to_string())?; - if is_io_source(value) { - let bytes = prefetched - .next() - .ok_or_else(|| "I/O-backed image source was not prefetched".to_string())?; - out.push(ImageSource::Bytes(bytes.to_vec())); - } else { - out.push(ImageSource::String(value.to_owned())); +) -> Result { + match item { + MmItem::Source(source) | MmItem::Ref { url: source } => { + if !is_io_source(&source) { + return Ok(ImageSource::String(source)); } - Ok(()) + prefetched + .next() + .map(|bytes| ImageSource::Bytes(bytes.to_vec())) + .ok_or_else(|| "I/O-backed image source was not prefetched".to_string()) } - Value::Binary(value) => { - out.push(ImageSource::Bytes(value.clone())); - Ok(()) - } - Value::Array(values) => { - for value in values { - match value { - Value::String(_) | Value::Binary(_) | Value::Nil => { - collect_images(value, prefetched, out)? - } - _ => { - return Err("unsupported image_data shape: nested/typed item".into()); - } - } - } - Ok(()) - } - _ => Err("unsupported image_data shape".into()), - } -} - -/// Rust mirror of Python `has_valid_data`: `nil` and (recursively) empty or -/// all-nil lists don't count as multimodal input. -pub fn value_present(value: &Value) -> bool { - match value { - Value::Nil => false, - Value::Array(values) => values.iter().any(value_present), - _ => true, + MmItem::Preprocessed { format } => Err(format!( + "unsupported image_data item: preprocessed `{format}` input" + )), } } @@ -127,65 +88,55 @@ pub fn value_present(value: &Value) -> bool { mod tests { use super::*; - fn image_work(image: Value) -> MmWorkItem { + fn src(s: &str) -> MmItem { + MmItem::Source(s.to_owned()) + } + + fn image_work(image_data: Vec) -> MmWorkItem { MmWorkItem { text: Some("prompt".into()), - image_data: Some(image), + image_data, ..Default::default() } } #[test] - fn converts_string_and_list_images() { - let one = to_mm_input(image_work(Value::from("data:image/png;base64,x"))).unwrap(); + fn converts_source_and_ref_images() { + let one = to_mm_input(image_work(vec![src("data:image/png;base64,x")])).unwrap(); assert_eq!(one.images.len(), 1); - let many = to_mm_input(image_work(Value::Array(vec![ - Value::from("a"), - Value::from("b"), - ]))) - .unwrap(); + let many = + to_mm_input(image_work(vec![src("a"), MmItem::Ref { url: "b".into() }])).unwrap(); assert_eq!(many.images.len(), 2); + assert!(matches!(&many.images[1], ImageSource::String(s) if s == "b")); } #[test] - fn unsupported_modalities_and_shapes_rejected() { + fn unsupported_modalities_and_items_rejected() { let video = MmWorkItem { - video_data: Some(Value::from("video.mp4")), + video_data: vec![src("video.mp4")], ..Default::default() }; assert!(to_mm_input(video).err().unwrap().contains("video/audio")); - let dict = Value::Map(vec![(Value::from("format"), Value::from("x"))]); - assert!( - to_mm_input(image_work(Value::Array(vec![dict]))) - .err() - .unwrap() - .contains("image_data shape") - ); - } - - #[test] - fn empty_video_audio_lists_are_not_modalities() { - // Mirrors Python `has_valid_data`: nil / empty lists don't count. - let work = MmWorkItem { - input_ids: Some(vec![1]), - image_data: Some(Value::from("a")), - video_data: Some(Value::Array(vec![])), - audio_data: Some(Value::Array(vec![Value::Array(vec![])])), - ..Default::default() - }; - assert_eq!(to_mm_input(work).unwrap().images.len(), 1); + let err = to_mm_input(image_work(vec![MmItem::Preprocessed { + format: "processor_output".into(), + }])) + .err() + .unwrap(); + assert!(err.contains("preprocessed `processor_output`"), "{err}"); } /// I/O-backed sources (URLs, file paths) take their prefetched bytes in walk /// order; one left unfetched errors, so no I/O can reach an MM worker. #[test] fn io_sources_use_prefetched_bytes() { - let image = Value::Array(vec![ - Value::from("http://a/x.png"), - Value::from("data:image/png;base64,x"), - Value::from("/mnt/nfs/y.png"), - ]); + let image = vec![ + src("http://a/x.png"), + src("data:image/png;base64,x"), + MmItem::Ref { + url: "/mnt/nfs/y.png".into(), + }, + ]; assert_eq!(io_sources(&image), vec!["http://a/x.png", "/mnt/nfs/y.png"]); let mut work = image_work(image.clone()); @@ -205,12 +156,6 @@ mod tests { #[test] fn image_free_work_rejected() { - assert!( - to_mm_input(image_work(Value::Nil)) - .err() - .unwrap() - .contains("no raw image sources") - ); assert!( to_mm_input(MmWorkItem::default()) .err() diff --git a/rust/sglang-server/src/tokenizer_manager/to_scheduler_tests.rs b/rust/sglang-server/src/tokenizer_manager/to_scheduler_tests.rs index c1b8139bc..dda406708 100644 --- a/rust/sglang-server/src/tokenizer_manager/to_scheduler_tests.rs +++ b/rust/sglang-server/src/tokenizer_manager/to_scheduler_tests.rs @@ -644,7 +644,9 @@ fn multimodal_sentinel_is_validated_after_expansion() { }; 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")), + audio_data: vec![crate::message::multimodal::MmItem::Source( + "data:audio/wav;base64,xxxx".into(), + )], ..Default::default() })); @@ -771,7 +773,9 @@ fn mm_generate_req(rid: &str) -> Request { rid: rid.to_string().into(), text: Some(" hi".into()), mm: Some(Box::new(crate::message::request::MmData { - image_data: Some(rmpv::Value::from("data:image/jpeg;base64,xxxx")), + image_data: vec![crate::message::multimodal::MmItem::Source( + "data:image/jpeg;base64,xxxx".into(), + )], ..Default::default() })), ..Default::default() @@ -826,7 +830,7 @@ fn mm_request_parks_then_mm_encoded_pushes_to_ring() { assert_eq!(sub.work.text.as_deref(), Some(" hi")); assert!(sub.work.input_ids.is_none(), "no client input_ids"); assert_eq!( - sub.work.image_data.as_ref().and_then(|v| v.as_str()), + sub.work.image_data.first().and_then(|item| item.source()), Some("data:image/jpeg;base64,xxxx") ); assert!(consumer.drain(16).headers.is_empty(), "parked, not queued");