diff --git a/Cargo.lock b/Cargo.lock index ad7ccd222d0e..8a8d7d80cbc6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3143,12 +3143,14 @@ dependencies = [ "anyhow", "async-stream", "async-trait", + "base64 0.22.1", "clap", "dynamo-backend-common", "dynamo-sidecar-common", "futures", "prost 0.13.5", "prost-types 0.13.5", + "serde", "serde_json", "tokio", "tokio-stream", diff --git a/lib/backend-common/src/lib.rs b/lib/backend-common/src/lib.rs index c7ef9b585435..0286d05d8f08 100644 --- a/lib/backend-common/src/lib.rs +++ b/lib/backend-common/src/lib.rs @@ -42,6 +42,7 @@ pub use engine::{ }; pub use error::{BackendError, DynamoError, ErrorType}; pub use metrics::{ComponentGauges, EngineMetrics, LifecycleGauges}; +pub use rl::RlWorkerMetadata; pub use run::{run, run_raw}; pub use snapshot_publisher::SnapshotPublisher; pub use worker::{RuntimeConfig, Worker, WorkerConfig}; diff --git a/lib/backend-common/src/rl.rs b/lib/backend-common/src/rl.rs index 9501c4a27040..9b87c481ae60 100644 --- a/lib/backend-common/src/rl.rs +++ b/lib/backend-common/src/rl.rs @@ -1,7 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -use std::sync::Arc; +use std::{num::NonZeroU32, sync::Arc}; use async_trait::async_trait; use dynamo_runtime::component::{Endpoint, StartedEndpoint}; @@ -24,6 +24,48 @@ pub(crate) struct RlServeEndpoint { pub(crate) struct RlEndpointConfig { endpoint_name: String, system_url: String, + metadata: Option, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct RlWorkerMetadata { + world_size: NonZeroU32, + weight_transfer_backend: Option, + admin_base_url: Option, +} + +impl RlWorkerMetadata { + pub fn new( + world_size: u32, + weight_transfer_backend: Option, + admin_base_url: Option, + ) -> anyhow::Result { + let world_size = NonZeroU32::new(world_size) + .ok_or_else(|| anyhow::anyhow!("RL worker world size must be positive"))?; + let weight_transfer_backend = normalized_optional( + weight_transfer_backend, + "RL weight-transfer backend must not be blank", + )?; + let admin_base_url = + normalized_optional(admin_base_url, "RL admin base URL must not be blank")?; + Ok(Self { + world_size, + weight_transfer_backend, + admin_base_url, + }) + } +} + +fn normalized_optional(value: Option, error: &str) -> anyhow::Result> { + value + .map(|value| { + let value = value.trim(); + if value.is_empty() { + anyhow::bail!(error.to_string()); + } + Ok(value.to_string()) + }) + .transpose() } impl RlServeEndpoint { @@ -32,7 +74,10 @@ impl RlServeEndpoint { } } -pub(crate) fn prepare_endpoint(primary: &Endpoint) -> anyhow::Result { +pub(crate) fn prepare_endpoint( + primary: &Endpoint, + metadata: Option, +) -> anyhow::Result { let endpoint_name = resolve_endpoint_name(&primary.id().name)?; let system_url = self_host_base_url(primary.drt()).ok_or_else(|| { anyhow::anyhow!( @@ -42,6 +87,7 @@ pub(crate) fn prepare_endpoint(primary: &Endpoint) -> anyhow::Result anyhow::Re struct RlRouteHandler { routes: EngineRouteRegistry, system_url: String, + metadata: Option, } impl RlRouteHandler { @@ -122,11 +170,21 @@ impl RlRouteHandler { let mut routes = self.routes.routes().into_iter().collect::>(); routes.sort(); routes.dedup(); - json!({ + let mut response = json!({ "status": "ok", "routes": routes, "system_url": self.system_url, - }) + }); + if let Some(metadata) = &self.metadata { + response["world_size"] = json!(metadata.world_size.get()); + if let Some(backend) = &metadata.weight_transfer_backend { + response["weight_transfer_backend"] = json!(backend); + } + if let Some(url) = &metadata.admin_base_url { + response["admin_base_url"] = json!(url); + } + } + response } } @@ -156,6 +214,14 @@ mod tests { let handler = RlRouteHandler { routes, system_url: "http://worker:8080".to_string(), + metadata: Some( + RlWorkerMetadata::new( + 4, + Some(" nccl ".to_string()), + Some(" http://worker:8120 ".to_string()), + ) + .expect("valid metadata"), + ), }; assert_eq!( @@ -164,6 +230,9 @@ mod tests { "status": "ok", "routes": ["control/pause_generation"], "system_url": "http://worker:8080", + "admin_base_url": "http://worker:8120", + "world_size": 4, + "weight_transfer_backend": "nccl", }) ); assert_eq!( @@ -171,4 +240,11 @@ mod tests { "error" ); } + + #[test] + fn rl_worker_metadata_rejects_invalid_values() { + assert!(RlWorkerMetadata::new(0, None, None).is_err()); + assert!(RlWorkerMetadata::new(1, Some(" ".to_string()), None).is_err()); + assert!(RlWorkerMetadata::new(1, None, Some(" ".to_string())).is_err()); + } } diff --git a/lib/backend-common/src/worker.rs b/lib/backend-common/src/worker.rs index 2a2fd940270b..d67cecc91622 100644 --- a/lib/backend-common/src/worker.rs +++ b/lib/backend-common/src/worker.rs @@ -193,6 +193,8 @@ pub struct WorkerConfig { pub route_to_encoder: bool, /// Publish the worker's engine routes through an auxiliary RL discovery endpoint. pub enable_rl: bool, + /// Optional RL topology and weight-transfer metadata published by the worker. + pub rl_metadata: Option, /// Optional frontend media decoding and fetch policy advertised on the /// model deployment card. pub media_decoder: Option, @@ -238,6 +240,7 @@ impl Default for WorkerConfig { runtime: RuntimeConfig::default(), route_to_encoder: false, enable_rl: false, + rl_metadata: None, media_decoder: None, media_fetcher: None, default_thinking_mode: None, @@ -926,12 +929,16 @@ impl Worker { let model_type = resolve_model_type(&self.config)?; let (worker_type, needs) = resolve_worker_type_and_needs(&self.config); let rl_config = if self.config.enable_rl { - Some(crate::rl::prepare_endpoint(&endpoint).map_err(|error| { - err( - ErrorType::Backend(BackendError::InvalidArgument), - format!("RL endpoint configuration: {error}"), - ) - })?) + Some( + crate::rl::prepare_endpoint(&endpoint, self.config.rl_metadata.clone()).map_err( + |error| { + err( + ErrorType::Backend(BackendError::InvalidArgument), + format!("RL endpoint configuration: {error}"), + ) + }, + )?, + ) } else { None }; diff --git a/lib/bindings/python/rust/backend.rs b/lib/bindings/python/rust/backend.rs index bdffba55b41d..03a92cbbf077 100644 --- a/lib/bindings/python/rust/backend.rs +++ b/lib/bindings/python/rust/backend.rs @@ -469,6 +469,7 @@ impl WorkerConfig { // Python vLLM owns and serves its existing `.rl` endpoint. // The shared Rust endpoint is opt-in for Rust sidecars only. enable_rl: false, + rl_metadata: None, media_decoder: media_decoder.map(|decoder| decoder.inner), media_fetcher: media_fetcher.map(|fetcher| fetcher.inner), }, diff --git a/lib/kv-router/src/protocols.rs b/lib/kv-router/src/protocols.rs index 4ea60c83df03..f7562b7c3ecb 100644 --- a/lib/kv-router/src/protocols.rs +++ b/lib/kv-router/src/protocols.rs @@ -79,6 +79,21 @@ pub fn pad_value_for_mm_hash(mm_hash: u64) -> u32 { (MM_PAD_SHIFT_VALUE + (mm_hash & MM_PAD_HASH_MASK)) as u32 } +/// Map a non-empty multimodal identifier to Dynamo's routing hash. +pub fn hash_mm_identifier(identifier: &str) -> Option { + if identifier.is_empty() { + return None; + } + if identifier.len() == 64 + && identifier + .chars() + .all(|character| character.is_ascii_hexdigit()) + { + return u64::from_str_radix(&identifier[..16], 16).ok(); + } + Some(xxh3::xxh3_64(identifier.as_bytes())) +} + /// Compute the hash for a sequence of tokens, optionally including multimodal metadata, /// LoRA adapter identity, and cache namespace. /// @@ -1865,6 +1880,18 @@ mod tests { ); } + #[test] + fn mm_identifier_hash_preserves_vllm_and_opaque_identifiers() { + let canonical = "0123456789abcdef".repeat(4); + assert_eq!(hash_mm_identifier(&canonical), Some(0x0123_4567_89ab_cdef)); + let opaque = "opaque-renderer-image-0"; + assert_eq!( + hash_mm_identifier(opaque), + Some(xxh3::xxh3_64(opaque.as_bytes())) + ); + assert_eq!(hash_mm_identifier(""), None); + } + #[test] fn test_router_event_new() { let worker_id = 0; diff --git a/lib/llm/src/discovery.rs b/lib/llm/src/discovery.rs index ae8dd5a6507c..5fb10953a263 100644 --- a/lib/llm/src/discovery.rs +++ b/lib/llm/src/discovery.rs @@ -2,7 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 mod model; -pub use model::Model; +pub use model::{GenerateEngineSelection, Model}; pub mod kv_source_membership; pub use kv_source_membership::{ diff --git a/lib/llm/src/discovery/model.rs b/lib/llm/src/discovery/model.rs index 839bc41d1158..795832bc1f5d 100644 --- a/lib/llm/src/discovery/model.rs +++ b/lib/llm/src/discovery/model.rs @@ -17,6 +17,7 @@ use super::worker_monitor::LoadThresholdConfig; use super::worker_set::WorkerSet; use crate::protocols::openai::ParsingOptions; +use crate::local_model::runtime_config::VLLM_EXACT_MM_ROUTING_CAPABILITY; use crate::types::{ RealtimeBidirectionalEngine, generic::tensor::TensorStreamingEngine, @@ -30,6 +31,13 @@ use crate::types::{ }, }; +#[derive(Clone)] +pub struct GenerateEngineSelection { + pub engine: GenerateStreamingEngine, + pub kv_cache_block_size: u32, + pub supports_exact_mm_routing: bool, +} + /// Emit a one-time deprecation warning when serving-readiness falls back to /// the legacy path because a namespace still contains a legacy card (a /// worker with no declared `worker_type`). Logged once per process to avoid @@ -671,6 +679,29 @@ impl Model { .ok_or_else(|| self.engine_error(self.has_generate_engine_for_capability(capability))) } + pub fn get_generate_engine_selection_for_capability( + &self, + capability: &str, + ) -> Result { + self.select_worker_set_with(|worker_set| { + worker_set + .supports_runtime_capability(capability) + .then(|| { + worker_set + .generate_engine + .clone() + .map(|engine| GenerateEngineSelection { + engine, + kv_cache_block_size: worker_set.card().kv_cache_block_size, + supports_exact_mm_routing: worker_set + .supports_runtime_capability(VLLM_EXACT_MM_ROUTING_CAPABILITY), + }) + }) + .flatten() + }) + .ok_or_else(|| self.engine_error(self.has_generate_engine_for_capability(capability))) + } + // -- Combined engine + parsing options (atomically from one WorkerSet) -- pub fn get_chat_engine_with_parsing( diff --git a/lib/llm/src/discovery/model_manager.rs b/lib/llm/src/discovery/model_manager.rs index 91f3ba0c1d03..320a74997127 100644 --- a/lib/llm/src/discovery/model_manager.rs +++ b/lib/llm/src/discovery/model_manager.rs @@ -20,7 +20,7 @@ use dynamo_kv_router::{ use super::worker_monitor::LoadThresholdConfig; use super::{ - KvSourceMembershipWatch, Model, RuntimeConfigWatch, WorkerSet, + GenerateEngineSelection, KvSourceMembershipWatch, Model, RuntimeConfigWatch, WorkerSet, kv_source_watch::KvSourceMembershipCoordinator, runtime_config_watch, }; @@ -1342,6 +1342,19 @@ impl ModelManager { .get_generate_engine_for_capability(capability) } + pub fn get_generate_engine_selection_for_capability( + &self, + model: &str, + capability: &str, + ) -> Result { + self.catalog + .load() + .models + .get(model) + .ok_or_else(|| ModelManagerError::ModelNotFound(model.to_string()))? + .get_generate_engine_selection_for_capability(capability) + } + // -- Combined engine + parsing options (atomically from one WorkerSet) -- pub fn get_chat_completions_engine_with_parsing( diff --git a/lib/llm/src/http/service/generate.rs b/lib/llm/src/http/service/generate.rs index 17f4bc8e1789..5b577013a495 100644 --- a/lib/llm/src/http/service/generate.rs +++ b/lib/llm/src/http/service/generate.rs @@ -32,8 +32,7 @@ use super::openai::{ }; use super::{RouteDoc, service_v2}; use crate::local_model::runtime_config::VLLM_INFERENCE_V1_GENERATE_CAPABILITY; -use crate::protocols::common::preprocessor::PreprocessedRequest; -use crate::protocols::common::{SamplingOptions, StopConditions}; +use crate::protocols::common::preprocessor::{MmRoutingInfo, PreprocessedRequest}; use crate::protocols::openai::generate::{ GenerateRequest, GenerateResponse, GenerateResponseOptions, SamplingParams, StreamOptions, }; @@ -173,6 +172,8 @@ struct VllmTitoEnvelope<'a> { priority: i32, #[serde(skip_serializing_if = "Option::is_none")] kv_transfer_params: Option<&'a serde_json::Map>, + #[serde(skip_serializing_if = "Option::is_none")] + features: Option<&'a crate::protocols::openai::generate::MultiModalFeatures>, #[serde(flatten)] passthrough: &'a serde_json::Map, } @@ -189,6 +190,7 @@ impl<'a> VllmTitoEnvelope<'a> { cache_salt, priority, kv_transfer_params, + features, passthrough, } = request; Self { @@ -200,46 +202,218 @@ impl<'a> VllmTitoEnvelope<'a> { cache_salt: cache_salt.as_deref(), priority: *priority, kv_transfer_params: kv_transfer_params.as_ref(), + features: features.as_ref(), passthrough, } } } +type MmPlaceholderRange = (usize, usize, u64, Option>); + +fn generate_mm_routing_info( + request: &GenerateRequest, + kv_cache_block_size: u32, +) -> Result, &'static str> { + let Some(features) = request.features.as_ref() else { + return Ok(None); + }; + if features + .mm_hashes + .keys() + .chain(features.mm_placeholders.keys()) + .any(|modality| modality != "image") + { + return Err( + "exact /generate multimodal routing currently supports image placeholders only", + ); + } + if kv_cache_block_size == 0 { + return Err("KV cache block size must be non-zero"); + } + let (Some(hashes), Some(placeholders)) = ( + features.mm_hashes.get("image"), + features.mm_placeholders.get("image"), + ) else { + return Ok(None); + }; + if hashes.len() != placeholders.len() { + return Err("image hashes and placeholders must have equal lengths"); + } + + let mut ranges: Vec = Vec::with_capacity(hashes.len()); + for (hash, placeholder) in hashes.iter().zip(placeholders) { + let hash = dynamo_kv_router::protocols::hash_mm_identifier(hash) + .ok_or("multimodal hashes must be non-empty strings")?; + let end = placeholder + .offset + .checked_add(placeholder.length) + .filter(|end| *end <= request.token_ids.len()) + .ok_or("multimodal placeholder range exceeds token_ids")?; + if placeholder.length == 0 { + return Err("multimodal placeholder lengths must be positive integers"); + } + if placeholder.is_embed.is_none() + && request.token_ids[placeholder.offset..end] + .windows(2) + .any(|pair| pair[0] != pair[1]) + { + return Err("mixed multimodal placeholder spans require is_embed"); + } + if placeholder + .is_embed + .as_ref() + .is_some_and(|mask| mask.len() != placeholder.length) + { + return Err("multimodal placeholder is_embed length must match placeholder length"); + } + ranges.push((placeholder.offset, end, hash, placeholder.is_embed.clone())); + } + if ranges.is_empty() { + return Ok(None); + } + + ranges.sort_unstable_by_key(|(offset, _, _, _)| *offset); + for pair in ranges.windows(2) { + let (_, previous_end, previous_hash, _) = &pair[0]; + let (next_offset, _, next_hash, _) = &pair[1]; + if previous_end > next_offset { + return Err("multimodal placeholder ranges must not overlap"); + } + if previous_end == next_offset && previous_hash != next_hash { + return Err("adjacent multimodal placeholders must share an identifier"); + } + } + + let block_size = kv_cache_block_size as usize; + for block_start in (0..request.token_ids.len()).step_by(block_size) { + let block_end = (block_start + block_size).min(request.token_ids.len()); + let mut worker_objects = Vec::new(); + let mut expected_by_position = vec![None; block_end - block_start]; + + for (offset, end, hash, is_embed) in &ranges { + let intersection_start = (*offset).max(block_start); + let intersection_end = (*end).min(block_end); + if intersection_start >= intersection_end { + continue; + } + worker_objects.push(*hash); + for global_position in intersection_start..intersection_end { + if is_embed + .as_ref() + .is_none_or(|mask| mask[global_position - *offset]) + { + expected_by_position[global_position - block_start] = Some(*hash); + } + } + } + + let mut expected_runs = Vec::new(); + let mut current_run = None; + for expected_hash in expected_by_position { + match (current_run, expected_hash) { + (None, Some(hash)) => { + current_run = Some(hash); + expected_runs.push(hash); + } + (Some(current), Some(hash)) if current != hash => { + return Err("adjacent multimodal embed positions must share an identifier"); + } + (Some(_), None) => current_run = None, + _ => {} + } + } + for (run_index, expected_hash) in expected_runs.into_iter().enumerate() { + let worker_hash = worker_objects + .get(run_index) + .or_else(|| worker_objects.last()) + .copied(); + if worker_hash != Some(expected_hash) { + return Err( + "sparse multimodal layout cannot be normalized exactly by worker events", + ); + } + } + } + + let mut routing_token_ids = request.token_ids.clone(); + for (offset, end, hash, is_embed) in ranges { + let pad = dynamo_kv_router::protocols::pad_value_for_mm_hash(hash); + if let Some(mask) = is_embed { + for (token, should_embed) in routing_token_ids[offset..end].iter_mut().zip(mask) { + if should_embed { + *token = pad; + } + } + } else { + routing_token_ids[offset..end].fill(pad); + } + } + let padded_len = routing_token_ids + .len() + .div_ceil(block_size) + .checked_mul(block_size) + .ok_or("multimodal routing token length overflow")?; + routing_token_ids.resize(padded_len, 0); + + Ok(Some(MmRoutingInfo { + routing_token_ids, + block_mm_infos: Vec::new(), + expanded_prompt_len: request.token_ids.len(), + })) +} + /// Project routing controls while retaining all engine-owned fields in /// `extra_args.vllm_tito`. The backend remains the authority for interpreting /// every vLLM-specific field. -fn preprocessed_from_generate( +fn preprocessed_from_generate_with_routing( request: GenerateRequest, model: &str, data_parallel_rank: Option, request_id: &str, + kv_cache_block_size: u32, + supports_exact_mm_routing: bool, ) -> anyhow::Result { let sampling = &request.sampling_params; let max_tokens = sampling.max_tokens(); - let min_tokens = sampling.min_tokens(); - let ignore_eos = sampling.ignore_eos(); + let stop_conditions = sampling.project_stop_conditions(); + let sampling_options = sampling.project_sampling_options(); + let output_options = sampling + .project_output_options() + .map_err(anyhow::Error::msg)?; + let cache_salt = match (request.cache_salt.as_deref(), sampling.cache_salt()) { + (Some(top_level), Some(sampling)) if top_level != sampling => { + anyhow::bail!("cache_salt conflicts with sampling_params.cache_salt"); + } + (Some(top_level), _) => Some(top_level.to_string()), + (None, Some(sampling)) => Some(sampling.to_string()), + (None, None) => None, + }; let routing_priority = dynamo_routing_priority(request.priority); + let mm_routing_info = if supports_exact_mm_routing { + match generate_mm_routing_info(&request, kv_cache_block_size) { + Ok(info) => info, + Err(reason) => { + tracing::debug!( + target: "mm_routing", + reason, + "invalid /generate multimodal routing metadata; using token-only routing" + ); + None + } + } + } else { + None + }; let vllm_tito = serde_json::to_value(VllmTitoEnvelope::new(&request, request_id))?; - let GenerateRequest { - token_ids, - cache_salt, - .. - } = request; + let GenerateRequest { token_ids, .. } = request; PreprocessedRequest::builder() .model(model.to_string()) .token_ids(token_ids) - .stop_conditions(StopConditions { - max_tokens, - min_tokens, - ignore_eos: Some(ignore_eos), - ..Default::default() - }) - .sampling_options(SamplingOptions { - n: Some(1), - ..Default::default() - }) - .output_options(Default::default()) + .stop_conditions(stop_conditions) + .sampling_options(sampling_options) + .output_options(output_options) + .mm_routing_info(mm_routing_info) .routing(Some(crate::protocols::common::preprocessor::RoutingHints { dp_rank: data_parallel_rank, expected_output_tokens: max_tokens, @@ -259,6 +433,23 @@ fn preprocessed_from_generate( .map_err(|error| anyhow::anyhow!("failed to build PreprocessedRequest: {error}")) } +#[cfg(test)] +fn preprocessed_from_generate( + request: GenerateRequest, + model: &str, + data_parallel_rank: Option, + request_id: &str, +) -> anyhow::Result { + preprocessed_from_generate_with_routing( + request, + model, + data_parallel_rank, + request_id, + 0, + false, + ) +} + /// Resolve, route, and dispatch a frontend-native token-in/token-out request. async fn handler_generate( State(state): State>, @@ -313,9 +504,9 @@ async fn handler_generate( return response.into_response(); } - let engine = match state + let engine_selection = match state .manager() - .get_generate_engine_for_capability(&model, VLLM_INFERENCE_V1_GENERATE_CAPABILITY) + .get_generate_engine_selection_for_capability(&model, VLLM_INFERENCE_V1_GENERATE_CAPABILITY) { Ok(engine) => engine, Err(error) => { @@ -330,11 +521,13 @@ async fn handler_generate( }; let request_context = resolve_generate_request_context(&headers, request.request_id.as_deref()); - let preprocessed = match preprocessed_from_generate( + let preprocessed = match preprocessed_from_generate_with_routing( request, &model, request_context.data_parallel_rank, &request_context.request_id, + engine_selection.kv_cache_block_size, + engine_selection.supports_exact_mm_routing, ) { Ok(preprocessed) => preprocessed, Err(error) => { @@ -371,7 +564,7 @@ async fn handler_generate( // each backend await point and then exits promptly. let response = match tokio::spawn( generate_dispatch( - engine, + engine_selection.engine, context, request_id, model, @@ -764,6 +957,7 @@ mod tests { let invalid = [ r#"{"token_ids":[1],"sampling_params":{},"stream_options":{"include_usage":true}}"#, r#"{"token_ids":[1],"sampling_params":{"max_tokens":0}}"#, + r#"{"token_ids":[1],"sampling_params":{"logprobs":-2}}"#, r#"{"token_ids":[1],"sampling_params":{"prompt_logprobs":-2}}"#, r#"{"token_ids":[1],"sampling_params":{"min_tokens":3,"max_tokens":2}}"#, ]; @@ -841,7 +1035,11 @@ mod tests { "stream": true, "stream_options": {"include_usage": true}, "cache_salt": "tenant-a", - "features": {"future_feature": [1, 2, 3]}, + "features": { + "mm_hashes": {"image": ["image-a"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 1}]}, + "kwargs_data": null + }, "priority": 7, "kv_transfer_params": {"remote": "worker-a"}, "future_top_level_field": {"anything": "works"} @@ -901,6 +1099,182 @@ mod tests { assert!(envelope.get("token_ids").is_none()); } + #[test] + fn distinct_images_with_identical_placeholders_get_distinct_routing_tokens() { + let request = |hash: &str| { + serde_json::from_value::(serde_json::json!({ + "token_ids": [10, 99, 99, 20], + "sampling_params": {}, + "features": { + "mm_hashes": {"image": [hash]}, + "mm_placeholders": {"image": [{"offset": 1, "length": 2}]}, + "kwargs_data": {"image": ["cmVk"]} + } + })) + .expect("deserialize multimodal request") + }; + let red = request("image-red"); + let blue = request("image-blue"); + + let red_routing = generate_mm_routing_info(&red, 4) + .expect("valid red routing") + .expect("red routing info"); + let blue_routing = generate_mm_routing_info(&blue, 4) + .expect("valid blue routing") + .expect("blue routing info"); + assert_ne!( + red_routing.routing_token_ids, + blue_routing.routing_token_ids + ); + assert_eq!(red.token_ids, blue.token_ids); + + let red_preprocessed = preprocessed_from_generate_with_routing( + red, + "test-model", + None, + "red-request", + 4, + true, + ) + .expect("build red request"); + assert_eq!( + red_preprocessed.extra_args.as_ref().unwrap()["vllm_tito"]["features"]["mm_hashes"]["image"] + [0], + "image-red" + ); + } + + #[test] + fn sampling_cache_salt_becomes_the_canonical_routing_namespace() { + let request: GenerateRequest = serde_json::from_value(serde_json::json!({ + "token_ids": [1, 2], + "sampling_params": { + "max_tokens": 8, + "cache_salt": "policy-version-3" + } + })) + .expect("deserialize request"); + + let preprocessed = + preprocessed_from_generate(request, "test-model", None, "resolved-request") + .expect("build request"); + + assert_eq!( + preprocessed + .routing + .as_ref() + .and_then(|routing| routing.cache_namespace.as_deref()), + Some("policy-version-3") + ); + } + + #[test] + fn conflicting_cache_salts_are_rejected_before_routing() { + let request: GenerateRequest = serde_json::from_value(serde_json::json!({ + "token_ids": [1, 2], + "cache_salt": "top-level", + "sampling_params": { + "max_tokens": 8, + "cache_salt": "sampling" + } + })) + .expect("deserialize request"); + + let error = preprocessed_from_generate(request, "test-model", None, "resolved-request") + .expect_err("reject conflicting cache salts"); + + assert_eq!( + error.to_string(), + "cache_salt conflicts with sampling_params.cache_salt" + ); + } + + #[test] + fn generate_projects_vllm_sampling_and_output_controls() { + let request: GenerateRequest = serde_json::from_value(serde_json::json!({ + "token_ids": [1, 2], + "sampling_params": { + "temperature": 0.25, + "top_p": 0.9, + "top_k": 17, + "min_p": 0.05, + "seed": 23, + "max_tokens": 8, + "min_tokens": 2, + "presence_penalty": 0.1, + "frequency_penalty": 0.2, + "repetition_penalty": 1.1, + "stop": ["done"], + "stop_token_ids": [7, 8], + "ignore_eos": true, + "logprobs": 3, + "prompt_logprobs": 4, + "skip_special_tokens": false, + "include_stop_str_in_output": true, + "return_token_ids": true + }, + "model": "test-model" + })) + .expect("deserialize request"); + + let preprocessed = + preprocessed_from_generate(request, "test-model", None, "resolved-request") + .expect("build request"); + + assert_eq!(preprocessed.sampling_options.temperature, Some(0.25)); + assert_eq!(preprocessed.sampling_options.top_p, Some(0.9)); + assert_eq!(preprocessed.sampling_options.top_k, Some(17)); + assert_eq!(preprocessed.sampling_options.min_p, Some(0.05)); + assert_eq!(preprocessed.sampling_options.seed, Some(23)); + assert_eq!(preprocessed.sampling_options.presence_penalty, Some(0.1)); + assert_eq!(preprocessed.sampling_options.frequency_penalty, Some(0.2)); + assert_eq!(preprocessed.sampling_options.repetition_penalty, Some(1.1)); + assert_eq!( + preprocessed.sampling_options.include_stop_str_in_output, + Some(true) + ); + assert_eq!(preprocessed.stop_conditions.max_tokens, Some(8)); + assert_eq!(preprocessed.stop_conditions.min_tokens, Some(2)); + assert_eq!( + preprocessed.stop_conditions.stop.as_deref(), + Some(&["done".to_string()][..]) + ); + assert_eq!( + preprocessed.stop_conditions.stop_token_ids.as_deref(), + Some(&[7, 8][..]) + ); + assert_eq!(preprocessed.stop_conditions.ignore_eos, Some(true)); + assert_eq!(preprocessed.output_options.logprobs, Some(3)); + assert_eq!(preprocessed.output_options.prompt_logprobs, Some(4)); + assert_eq!(preprocessed.output_options.skip_special_tokens, Some(false)); + assert_eq!( + preprocessed.output_options.return_tokens_as_token_ids, + Some(true) + ); + } + + #[test] + fn generate_projects_all_logprobs_and_unbounded_top_k() { + let request: GenerateRequest = serde_json::from_value(serde_json::json!({ + "token_ids": [1, 2], + "sampling_params": { + "top_k": -1, + "logprobs": -1, + "prompt_logprobs": -1 + }, + "model": "test-model" + })) + .expect("deserialize request"); + + let preprocessed = + preprocessed_from_generate(request, "test-model", None, "resolved-request") + .expect("build request"); + + assert_eq!(preprocessed.sampling_options.top_k, Some(-1)); + assert_eq!(preprocessed.output_options.logprobs, Some(u32::MAX)); + assert_eq!(preprocessed.output_options.prompt_logprobs, Some(u32::MAX)); + } + #[test] fn omitted_max_tokens_stays_omitted_in_control_shadow() { let request: GenerateRequest = serde_json::from_value(serde_json::json!({ diff --git a/lib/llm/src/local_model/runtime_config.rs b/lib/llm/src/local_model/runtime_config.rs index 9e6a3ca4d6f6..9e45e623bdd1 100644 --- a/lib/llm/src/local_model/runtime_config.rs +++ b/lib/llm/src/local_model/runtime_config.rs @@ -80,6 +80,9 @@ pub const ENV_TOKENIZER_FALLBACK: &str = "DYN_TOKENIZER_FALLBACK"; /// surfaces without implementing vLLM's Generate contract. pub const VLLM_INFERENCE_V1_GENERATE_CAPABILITY: &str = "vllm_inference_v1_generate"; +/// Worker-advertised support for vLLM-compatible multimodal KV-event hashing. +pub const VLLM_EXACT_MM_ROUTING_CAPABILITY: &str = "vllm_exact_mm_routing"; + /// Worker-advertised support for Dynamo's SGLang-compatible `POST /generate` /// adapter. /// diff --git a/lib/llm/src/protocols/openai/generate.rs b/lib/llm/src/protocols/openai/generate.rs index 660391e43cb4..8af8212edb9f 100644 --- a/lib/llm/src/protocols/openai/generate.rs +++ b/lib/llm/src/protocols/openai/generate.rs @@ -10,7 +10,7 @@ //! in `passthrough`. `sampling_params` is validated while its complete JSON //! object is retained for the version-matched worker. -use std::collections::HashMap; +use std::collections::{BTreeMap, HashMap}; use anyhow::Result; use dynamo_runtime::error::{BackendError, DynamoError, ErrorType as DynamoErrorType}; @@ -22,6 +22,7 @@ use serde_json::{Map, Value}; use super::{convert_backend_top_logprobs, token_to_utf8_bytes}; use crate::protocols::Annotated; use crate::protocols::common::llm_backend::{LLMEngineOutput, PromptLogprobs}; +use crate::protocols::common::{OutputOptions, SamplingOptions, StopConditions}; /// Token-in/token-out generation request. /// @@ -64,8 +65,11 @@ pub struct GenerateRequest { #[serde(default, skip_serializing_if = "Option::is_none")] pub kv_transfer_params: Option>, - /// Future top-level fields, including Python-frontend-only fields such as - /// `features`, are retained and forwarded to the worker. + /// Engine-ready multimodal features emitted by vLLM's render endpoint. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub features: Option, + + /// Future top-level fields are retained and forwarded to the worker. #[serde(flatten)] pub passthrough: Map, } @@ -92,6 +96,13 @@ impl GenerateRequest { return Err("sampling_params.max_tokens must be greater than 0.".to_string()); } + if let Some(logprobs) = self.sampling_params.logprobs() + && logprobs < 0 + && logprobs != -1 + { + return Err("sampling_params.logprobs must be non-negative or -1.".to_string()); + } + if let Some(prompt_logprobs) = self.sampling_params.prompt_logprobs() { if prompt_logprobs < 0 && prompt_logprobs != -1 { return Err( @@ -116,10 +127,101 @@ impl GenerateRequest { )); } + if let Some(features) = &self.features { + features.validate(self.token_ids.len())?; + } + Ok(()) } } +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] +pub struct MultiModalFeatures { + pub mm_hashes: BTreeMap>, + pub mm_placeholders: BTreeMap>, + #[serde(default)] + pub kwargs_data: Option>>>, +} + +impl MultiModalFeatures { + fn validate(&self, token_count: usize) -> Result<(), String> { + if self.mm_hashes.is_empty() { + return Err("features.mm_hashes must contain at least one modality.".to_string()); + } + if self.mm_hashes.keys().ne(self.mm_placeholders.keys()) { + return Err( + "features.mm_hashes and features.mm_placeholders must contain the same modalities." + .to_string(), + ); + } + if let Some(kwargs_data) = &self.kwargs_data + && self.mm_hashes.keys().ne(kwargs_data.keys()) + { + return Err( + "features.kwargs_data must contain the same modalities as features.mm_hashes." + .to_string(), + ); + } + + let mut ranges = Vec::new(); + for (modality, hashes) in &self.mm_hashes { + let placeholders = &self.mm_placeholders[modality]; + if hashes.len() != placeholders.len() { + return Err(format!( + "features.{modality} hashes and placeholders must have equal lengths." + )); + } + if let Some(kwargs_data) = &self.kwargs_data + && kwargs_data[modality].len() != hashes.len() + { + return Err(format!( + "features.kwargs_data.{modality} must align with hashes and placeholders." + )); + } + for (index, (hash, placeholder)) in hashes.iter().zip(placeholders).enumerate() { + if hash.is_empty() { + return Err(format!( + "features.mm_hashes.{modality}[{index}] must be non-empty." + )); + } + if placeholder.length == 0 { + return Err(format!( + "features.mm_placeholders.{modality}[{index}].length must be positive." + )); + } + let end = placeholder + .offset + .checked_add(placeholder.length) + .filter(|end| *end <= token_count) + .ok_or_else(|| { + format!("features.mm_placeholders.{modality}[{index}] exceeds token_ids.") + })?; + if let Some(mask) = &placeholder.is_embed + && mask.len() != placeholder.length + { + return Err(format!( + "features.mm_placeholders.{modality}[{index}].is_embed must match length." + )); + } + ranges.push((placeholder.offset, end)); + } + } + ranges.sort_unstable(); + if ranges.windows(2).any(|pair| pair[0].1 > pair[1].0) { + return Err("features multimodal placeholder ranges must not overlap.".to_string()); + } + Ok(()) + } +} + +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] +pub struct MultiModalPlaceholderRange { + pub offset: usize, + pub length: usize, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub is_embed: Option>, +} + /// vLLM Rust streaming options. Unknown options are retained for forward /// compatibility even though the current unary Dynamo profile does not consume /// them. @@ -151,7 +253,7 @@ pub struct SamplingParams { // reads only the controls it needs; `raw` remains authoritative. temperature: Option, top_p: Option, - top_k: Option, + top_k: Option, seed: Option, max_tokens: Option, min_tokens: Option, @@ -162,8 +264,12 @@ pub struct SamplingParams { frequency_penalty: Option, presence_penalty: Option, repetition_penalty: Option, + stop: Option, stop_token_ids: Option>, ignore_eos: bool, + skip_special_tokens: Option, + include_stop_str_in_output: Option, + return_token_ids: Option, logit_bias: Option>, allowed_token_ids: Option>, bad_words: Option>, @@ -172,9 +278,26 @@ pub struct SamplingParams { /// typed view opaque avoids duplicating version-specific vLLM validation. structured_outputs: Option, skip_reading_prefix_cache: Option, + cache_salt: Option, vllm_xargs: Option>, } +#[derive(Debug, Clone, Deserialize)] +#[serde(untagged)] +enum StopStrings { + One(String), + Many(Vec), +} + +impl StopStrings { + fn as_vec(&self) -> Vec { + match self { + Self::One(value) => vec![value.clone()], + Self::Many(values) => values.clone(), + } + } +} + impl SamplingParams { pub fn max_tokens(&self) -> Option { self.max_tokens @@ -196,9 +319,61 @@ impl SamplingParams { self.prompt_logprobs } + pub fn cache_salt(&self) -> Option<&str> { + self.cache_salt.as_deref() + } + pub fn as_value(&self) -> &Value { &self.raw } + + pub(crate) fn project_sampling_options(&self) -> SamplingOptions { + SamplingOptions { + n: Some(1), + presence_penalty: self.presence_penalty, + frequency_penalty: self.frequency_penalty, + repetition_penalty: self.repetition_penalty, + temperature: self.temperature, + top_p: self.top_p, + top_k: self.top_k, + min_p: self.min_p, + seed: self.seed, + include_stop_str_in_output: self.include_stop_str_in_output, + ..Default::default() + } + } + + pub(crate) fn project_stop_conditions(&self) -> StopConditions { + StopConditions { + max_tokens: self.max_tokens, + stop: self.stop.as_ref().map(StopStrings::as_vec), + stop_token_ids: self.stop_token_ids.clone(), + min_tokens: self.min_tokens, + ignore_eos: Some(self.ignore_eos), + ..Default::default() + } + } + + pub(crate) fn project_output_options(&self) -> Result { + Ok(OutputOptions { + logprobs: project_logprob_count(self.logprobs, "logprobs")?, + prompt_logprobs: project_logprob_count(self.prompt_logprobs, "prompt_logprobs")?, + skip_special_tokens: self.skip_special_tokens, + return_tokens_as_token_ids: self.return_token_ids, + ..Default::default() + }) + } +} + +fn project_logprob_count(value: Option, field: &str) -> Result, String> { + match value { + None => Ok(None), + Some(-1) => Ok(Some(u32::MAX)), + Some(value) if value >= 0 => Ok(Some(value as u32)), + Some(value) => Err(format!( + "sampling_params.{field} must be non-negative or -1, got {value}" + )), + } } impl Serialize for SamplingParams { @@ -245,15 +420,20 @@ impl<'de> Deserialize<'de> for SamplingParams { frequency_penalty: field!(frequency_penalty), presence_penalty: field!(presence_penalty), repetition_penalty: field!(repetition_penalty), + stop: field!(stop), stop_token_ids: field!(stop_token_ids), ignore_eos: sampling_field_or_default(object, "ignore_eos") .map_err(serde::de::Error::custom)?, + skip_special_tokens: field!(skip_special_tokens), + include_stop_str_in_output: field!(include_stop_str_in_output), + return_token_ids: field!(return_token_ids), logit_bias: field!(logit_bias), allowed_token_ids: field!(allowed_token_ids), bad_words: field!(bad_words), logprob_token_ids: field!(logprob_token_ids), structured_outputs: field!(structured_outputs), skip_reading_prefix_cache: field!(skip_reading_prefix_cache), + cache_salt: field!(cache_salt), vllm_xargs: field!(vllm_xargs), raw, }) @@ -295,7 +475,7 @@ pub struct GenerateResponseChoice { pub finish_reason: Option, - pub routed_experts: Option, + pub routed_experts: Option, } /// Token-in/token-out generation response. @@ -323,7 +503,7 @@ struct GenerateChoiceAcc { token_ids: Vec, logprobs: Option>, finish_reason: Option, - routed_experts: Option, + routed_experts: Option, } impl GenerateChoiceAcc { @@ -547,11 +727,7 @@ impl GenerateAggregator { }); if let Some(engine_data) = output.engine_data.as_ref() { if let Some(routed_experts) = engine_data.get("routed_experts") { - choice.routed_experts = Some( - serde_json::from_value(routed_experts.clone()).map_err(|error| { - anyhow::anyhow!("invalid generate routed_experts payload: {error}") - })?, - ); + choice.routed_experts = Some(routed_experts.clone()); } if let Some(kv_transfer_params) = engine_data.get("kv_transfer_params") { self.kv_transfer_params = Some(kv_transfer_params.clone()); @@ -690,6 +866,51 @@ mod tests { assert_eq!(back.get("future_field"), Some(&json!("kept"))); } + #[test] + fn generate_request_types_and_validates_multimodal_features() { + let req: GenerateRequest = serde_json::from_value(json!({ + "token_ids": [10, 99, 99, 20], + "sampling_params": {}, + "features": { + "mm_hashes": {"image": ["image-a"]}, + "mm_placeholders": {"image": [{"offset": 1, "length": 2}]}, + "kwargs_data": {"image": ["encoded-kwargs"]} + } + })) + .expect("deserialize multimodal request"); + + assert!(req.validate().is_ok()); + assert!(!req.passthrough.contains_key("features")); + let features = req.features.expect("typed features"); + assert_eq!(features.mm_hashes["image"], ["image-a"]); + assert_eq!(features.mm_placeholders["image"][0].offset, 1); + assert_eq!( + features.kwargs_data.expect("kwargs data")["image"], + [Some("encoded-kwargs".to_string())] + ); + } + + #[test] + fn generate_request_rejects_overlapping_multimodal_features() { + let req: GenerateRequest = serde_json::from_value(json!({ + "token_ids": [10, 99, 99, 99, 20], + "sampling_params": {}, + "features": { + "mm_hashes": {"image": ["image-a", "image-b"]}, + "mm_placeholders": {"image": [ + {"offset": 1, "length": 2}, + {"offset": 2, "length": 2} + ]} + } + })) + .expect("deserialize multimodal request"); + + assert_eq!( + req.validate().expect_err("overlap must fail"), + "features multimodal placeholder ranges must not overlap." + ); + } + #[test] fn generate_request_preserves_unknown_sampling_fields() { let raw_sampling = json!({ @@ -732,10 +953,6 @@ mod tests { "token_ids": [1], "sampling_params": {"max_tokens": -1} }), - json!({ - "token_ids": [1], - "sampling_params": {"top_k": -1} - }), json!({ "token_ids": [1], "sampling_params": {"ignore_eos": null} @@ -743,6 +960,13 @@ mod tests { ] { assert!(serde_json::from_value::(raw).is_err()); } + + let unbounded_top_k: GenerateRequest = serde_json::from_value(json!({ + "token_ids": [1], + "sampling_params": {"top_k": -1} + })) + .expect("vLLM uses top_k=-1 to disable top-k filtering"); + assert_eq!(unbounded_top_k.sampling_params.top_k, Some(-1)); } #[test] @@ -1041,8 +1265,8 @@ mod tests { assert_eq!(response.choices[0].token_ids, Some(vec![100, 101])); assert_eq!(response.choices[0].finish_reason.as_deref(), Some("length")); assert_eq!( - response.choices[0].routed_experts.as_deref(), - Some("encoded-experts") + response.choices[0].routed_experts, + Some(json!("encoded-experts")) ); let logprobs = response.choices[0] .logprobs @@ -1105,35 +1329,40 @@ mod tests { .expect("aggregate routed experts"); assert_eq!(response.choices[0].index, 0); - assert_eq!( - response.choices[0].routed_experts.as_deref(), - Some("experts-0") - ); + assert_eq!(response.choices[0].routed_experts, Some(json!("experts-0"))); assert_eq!(response.choices[1].index, 1); - assert_eq!( - response.choices[1].routed_experts.as_deref(), - Some("experts-1") - ); + assert_eq!(response.choices[1].routed_experts, Some(json!("experts-1"))); } #[tokio::test] - async fn generate_response_rejects_malformed_routed_experts() { + async fn generate_response_preserves_structured_routed_experts() { let stream = futures::stream::iter([Annotated::from_data(LLMEngineOutput { token_ids: vec![100], index: Some(0), finish_reason: Some(crate::protocols::common::FinishReason::Stop), - engine_data: Some(json!({"routed_experts": {"unexpected": "object"}})), + engine_data: Some(json!({ + "routed_experts": { + "data": "AQIDBA==", + "shape": [2, 1, 2], + "start": 3, + "dtype": "uint8" + } + })), ..Default::default() })]); - let error = GenerateResponse::from_annotated_stream(stream, "req-routed".to_string()) + let response = GenerateResponse::from_annotated_stream(stream, "req-routed".to_string()) .await - .expect_err("malformed routed experts must fail"); + .expect("structured routed experts must pass through"); - assert!( - error - .to_string() - .contains("invalid generate routed_experts payload") + assert_eq!( + response.choices[0].routed_experts, + Some(json!({ + "data": "AQIDBA==", + "shape": [2, 1, 2], + "start": 3, + "dtype": "uint8" + })) ); } diff --git a/lib/rl/src/lib.rs b/lib/rl/src/lib.rs index 15b0626ca976..4f5086935295 100644 --- a/lib/rl/src/lib.rs +++ b/lib/rl/src/lib.rs @@ -40,6 +40,7 @@ const DEFAULT_REQUEST_TIMEOUT_SECS: u64 = 30; /// Global cap on concurrent per-worker probes (across all in-flight discovery /// requests), so a large fleet or many concurrent callers can't fan out without bound. const DEFAULT_MAX_CONCURRENT_PROBES: usize = 32; +const RL_WORKERS_PROTOCOL_VERSION: u32 = 1; type ModelKey = (String, String, u64); @@ -104,14 +105,21 @@ pub struct RlWorkerInfo { #[serde(skip_serializing_if = "Option::is_none")] pub system_url: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub admin_base_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub model: Option, pub routes: Vec, #[serde(skip_serializing_if = "Option::is_none")] + pub world_size: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub weight_transfer_backend: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub error: Option, } #[derive(Debug, Clone, serde::Serialize)] pub struct RlWorkersResponse { + pub protocol_version: u32, pub namespace: String, pub workers: Vec, } @@ -195,6 +203,7 @@ pub fn rl_router(state: RlDiscoveryState) -> Router { async fn workers_handler(State(state): State) -> impl IntoResponse { match list_workers(&state).await { Ok(workers) => Json(RlWorkersResponse { + protocol_version: RL_WORKERS_PROTOCOL_VERSION, namespace: state.config.namespace.clone(), workers, }) @@ -229,8 +238,7 @@ async fn list_workers(state: &RlDiscoveryState) -> anyhow::Result anyhow::Result>(); // Bound the client cache (N2): drop clients for endpoints that are no longer @@ -315,13 +330,34 @@ async fn describe_worker( call_worker_routes(state, &endpoint, timeout).await }; match tokio::time::timeout(timeout, probe).await { - Ok(Ok(routes)) => worker_info(endpoint, model, routes.routes, routes.system_url, None), - Ok(Err(err)) => worker_info(endpoint, model, Vec::new(), None, Some(err.to_string())), + Ok(Ok(routes)) => worker_info( + endpoint, + model, + routes.routes, + routes.system_url, + routes.admin_base_url, + routes.world_size, + routes.weight_transfer_backend, + None, + ), + Ok(Err(err)) => worker_info( + endpoint, + model, + Vec::new(), + None, + None, + None, + None, + Some(err.to_string()), + ), Err(_) => worker_info( endpoint, model, Vec::new(), None, + None, + None, + None, Some(format!( "worker discovery timed out after {}s", timeout.as_secs() @@ -334,6 +370,9 @@ async fn describe_worker( struct WorkerRoutes { routes: Vec, system_url: Option, + admin_base_url: Option, + world_size: Option, + weight_transfer_backend: Option, } async fn call_worker_routes( @@ -440,7 +479,46 @@ fn parse_worker_routes(value: serde_json::Value) -> anyhow::Result .filter(|url| !url.is_empty()) .map(ToString::to_string); - Ok(WorkerRoutes { routes, system_url }) + let admin_base_url = value + .get("admin_base_url") + .and_then(|url| url.as_str()) + .map(str::trim) + .filter(|url| !url.is_empty()) + .map(ToString::to_string); + + let world_size = value + .get("world_size") + .map(|value| { + let value = value.as_u64().ok_or_else(|| { + anyhow::anyhow!("worker routes response has invalid 'world_size'") + })?; + u32::try_from(value) + .ok() + .filter(|value| *value > 0) + .ok_or_else(|| anyhow::anyhow!("worker routes response has invalid 'world_size'")) + }) + .transpose()?; + let weight_transfer_backend = value + .get("weight_transfer_backend") + .map(|value| { + let value = value.as_str().ok_or_else(|| { + anyhow::anyhow!("worker routes response has invalid 'weight_transfer_backend'") + })?; + let value = value.trim(); + if value.is_empty() { + anyhow::bail!("worker routes response has invalid 'weight_transfer_backend'"); + } + Ok(value.to_string()) + }) + .transpose()?; + + Ok(WorkerRoutes { + routes, + system_url, + admin_base_url, + world_size, + weight_transfer_backend, + }) } fn worker_info( @@ -448,6 +526,9 @@ fn worker_info( model: Option, mut routes: Vec, system_url: Option, + admin_base_url: Option, + world_size: Option, + weight_transfer_backend: Option, error: Option, ) -> RlWorkerInfo { routes.sort(); @@ -461,8 +542,11 @@ fn worker_info( instance_id: endpoint.instance_id, transport: endpoint.transport, system_url, + admin_base_url, model, routes, + world_size, + weight_transfer_backend, error, } } @@ -537,12 +621,18 @@ mod tests { let parsed = parse_worker_routes(json!({ "routes": ["pause_generation", "resume_generation"], "system_url": " http://worker:8080 ", + "admin_base_url": " http://worker:8120 ", + "world_size": 4, + "weight_transfer_backend": "nccl", })) .expect("valid payload"); let routes: Vec<&str> = parsed.routes.iter().map(String::as_str).collect(); assert_eq!(routes, ["pause_generation", "resume_generation"]); // system_url is trimmed. assert_eq!(parsed.system_url.as_deref(), Some("http://worker:8080")); + assert_eq!(parsed.admin_base_url.as_deref(), Some("http://worker:8120")); + assert_eq!(parsed.world_size, Some(4)); + assert_eq!(parsed.weight_transfer_backend.as_deref(), Some("nccl")); } #[test] @@ -572,6 +662,18 @@ mod tests { assert!(err.to_string().contains("empty route entry")); } + #[test] + fn parse_worker_routes_rejects_invalid_rl_metadata() { + let zero = parse_worker_routes(json!({ "routes": [], "world_size": 0 })).unwrap_err(); + assert!(zero.to_string().contains("world_size")); + let blank = parse_worker_routes(json!({ + "routes": [], + "weight_transfer_backend": " ", + })) + .unwrap_err(); + assert!(blank.to_string().contains("weight_transfer_backend")); + } + #[test] fn parse_worker_routes_propagates_worker_error_status() { let err = parse_worker_routes(json!({ "status": "error", "message": "engine is dead" })) diff --git a/lib/sidecar/vllm/Cargo.toml b/lib/sidecar/vllm/Cargo.toml index ea6603f2714f..948d30e62830 100644 --- a/lib/sidecar/vllm/Cargo.toml +++ b/lib/sidecar/vllm/Cargo.toml @@ -26,8 +26,10 @@ dynamo-sidecar-common = { workspace = true } anyhow = { workspace = true } async-stream = { workspace = true } async-trait = { workspace = true } +base64 = "0.22" clap = { version = "4", features = ["derive", "env"] } futures = { workspace = true } +serde = { workspace = true } serde_json = { workspace = true } tokio = { workspace = true } tokio-util = { workspace = true } diff --git a/lib/sidecar/vllm/proto/README.md b/lib/sidecar/vllm/proto/README.md index 4dabfc683960..258e6c1af9d4 100644 --- a/lib/sidecar/vllm/proto/README.md +++ b/lib/sidecar/vllm/proto/README.md @@ -5,9 +5,9 @@ SPDX-License-Identifier: Apache-2.0 # Vendored vLLM protocol -- Inference source: [`rust/proto/inference.proto`](https://github.com/vllm-project/vllm/blob/3d1f5cee1552b8208f3009c75f8bc856f27e0eff/rust/proto/inference.proto) at `3d1f5cee1552b8208f3009c75f8bc856f27e0eff` +- Inference source: [`rust/proto/inference.proto`](https://github.com/vllm-project/vllm/blob/5fd7a888386cff800f32de6b5a33d1dd3ca1e397/rust/proto/inference.proto) at `5fd7a888386cff800f32de6b5a33d1dd3ca1e397` - RL Control source: [`rust/proto/control.proto`](https://github.com/vllm-project/vllm/blob/76ebe5a217d7536a5661272c680f0b1e3a62f5be/rust/proto/control.proto) from [vllm-project/vllm#51316](https://github.com/vllm-project/vllm/pull/51316) at `76ebe5a217d7536a5661272c680f0b1e3a62f5be` -- `inference.proto` SHA-256: `6152c306583166ecd691c9c715cab950523e8d1ed2db3dc2bcb538f6ca90e56f` +- `inference.proto` SHA-256: `4c04f91d4967d1ba873fff6f546df138bc15cd29565c707c8554163392bb609a` - `control.proto` SHA-256: `db72b0782142054293b07fd48247cc821c048213b9c95dbc37fb0d81dde8f46f` The files are copied without modification. Update the revision and checksums together. `dynamo-vllm-sidecar` generates and temporarily exports these types for `dynamo-vllm-mocker-server`. diff --git a/lib/sidecar/vllm/proto/inference.proto b/lib/sidecar/vllm/proto/inference.proto index 021d93a1f7fb..f9ca2dd68f06 100644 --- a/lib/sidecar/vllm/proto/inference.proto +++ b/lib/sidecar/vllm/proto/inference.proto @@ -107,6 +107,10 @@ message ResponseOptions { bool output_token_ids = 5; bool output_logprobs = 6; optional CandidateTokens output_candidates = 7; + // Defaults to true when omitted. + optional bool skip_special_tokens = 8; + // Skip routing rows for this already-returned prompt prefix. + optional uint32 routed_experts_prompt_start = 9; } message KVCacheParameters { @@ -152,6 +156,15 @@ message SequenceOutput { // Only present in final output for this sequence optional FinishInfo finish_info = 8; + // Only present in final output when routed-experts capture is enabled. + optional RoutedExperts routed_experts = 9; +} + +message RoutedExperts { + bytes data = 1; + repeated uint32 shape = 2; + string dtype = 3; + uint32 start = 4; } // Prompt info, returned in the first response @@ -223,7 +236,19 @@ message MediaItem { string url = 2; // http:// or https:// string data_uri = 3; // data: bytes raw_bytes = 4; + PreprocessedMediaFeatures features = 7; } string mime_type = 5; string uuid = 6; } + +// Engine-ready multimodal features produced by a version-compatible frontend. +// `kwargs` is the MessagePack wire representation of one MmKwargsItem. +message PreprocessedMediaFeatures { + optional bytes kwargs = 1; + string identifier = 2; + uint64 offset = 3; + uint64 length = 4; + optional string mm_hash = 5; + repeated bool is_embed = 6; +} diff --git a/lib/sidecar/vllm/src/args.rs b/lib/sidecar/vllm/src/args.rs index 2dc8a67ee3bb..4385bc0268bd 100644 --- a/lib/sidecar/vllm/src/args.rs +++ b/lib/sidecar/vllm/src/args.rs @@ -15,4 +15,8 @@ pub(crate) struct Args { /// vLLM gRPC endpoint as host:port or an http:// URL. #[arg(long, env = "VLLM_GRPC_ENDPOINT")] pub vllm_endpoint: String, + + /// Optional vLLM-RS HTTP endpoint used for development collective RPCs. + #[arg(long, env = "VLLM_HTTP_ENDPOINT")] + pub vllm_http_endpoint: Option, } diff --git a/lib/sidecar/vllm/src/convert.rs b/lib/sidecar/vllm/src/convert.rs index 9f93ca4e2abe..eb9fde1bc704 100644 --- a/lib/sidecar/vllm/src/convert.rs +++ b/lib/sidecar/vllm/src/convert.rs @@ -1,6 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; use dynamo_backend_common::{ DisaggregationMode, DynamoError, GuidedDecodingOptions, LLMEngineOutput, MultimodalData, PrefillResult, PreprocessedRequest, StopReason, TopLogprob, usage, @@ -16,6 +17,29 @@ const MM_HASHES_KEY: &str = "mm_hashes"; // Must match DYNAMO_CACHE_SALT_PREFIX in lib/kv-router/src/zmq_wire/extra_keys.rs. const DYNAMO_CACHE_SALT_PREFIX: &str = "dynamo-cache-salt:"; +#[derive(Debug, serde::Deserialize)] +struct VllmTitoFeatures { + mm_hashes: std::collections::BTreeMap>, + mm_placeholders: std::collections::BTreeMap>, + #[serde(default)] + kwargs_data: Option>>>, +} + +#[derive(Debug, serde::Deserialize)] +struct VllmTitoPlaceholder { + offset: u64, + length: u64, + #[serde(default)] + is_embed: Option>, +} + +#[derive(Default)] +struct VllmTitoProjection { + priority: Option, + features: Option, + routed_experts_prompt_start: Option, +} + pub(crate) fn build_generate_request( request: PreprocessedRequest, request_id: String, @@ -23,24 +47,24 @@ pub(crate) fn build_generate_request( ) -> Result { validate_request(&request, mode)?; - let has_media = request + let has_raw_media = request .multi_modal_data .as_ref() .is_some_and(|media| media.values().any(|items| !items.is_empty())); // Decode reuses the prefill-expanded tokens without reprocessing media. - let forwarded_mm_uuids = if has_media && !mode.is_decode() { + let forwarded_mm_uuids = if has_raw_media && !mode.is_decode() { forwarded_mm_uuids(&request)? } else { None }; - let media = if mode.is_decode() { + let raw_media = if mode.is_decode() { Vec::new() } else { build_media(&request, forwarded_mm_uuids.as_deref())? }; let mut prefill_result = request.prefill_result; let mut token_ids = request.token_ids; - if mode.is_decode() && has_media { + if mode.is_decode() && has_raw_media { token_ids = take_multimodal_prompt_token_ids(&mut prefill_result)?; } let prompt_logprobs = request.output_options.prompt_logprobs; @@ -56,7 +80,7 @@ pub(crate) fn build_generate_request( request.stop_conditions.min_tokens.unwrap_or(0) }; let mut routing = request.routing; - let priority = routing + let mut priority = routing .as_ref() .and_then(|routing| routing.priority) .unwrap_or(0); @@ -67,6 +91,35 @@ pub(crate) fn build_generate_request( let sampling = request.sampling_options; let stop_conditions = request.stop_conditions; let mut extra_args = request.extra_args; + let vllm_tito = validate_and_remove_vllm_tito( + &mut extra_args, + cache_salt.as_deref(), + priority, + &request_id, + )?; + if let Some(vllm_priority) = vllm_tito.priority { + priority = vllm_priority; + } + let feature_media = if mode.is_decode() { + Vec::new() + } else { + vllm_tito + .features + .map(build_preprocessed_media) + .transpose()? + .unwrap_or_default() + }; + if !raw_media.is_empty() && !feature_media.is_empty() { + return Err(client::invalid_argument( + "raw multimodal data and preprocessed features cannot be mixed", + )); + } + let media = if feature_media.is_empty() { + raw_media + } else { + feature_media + }; + let has_media = has_raw_media || !media.is_empty(); consume_redundant_nvext(&mut extra_args, cache_salt.as_deref())?; if has_media && let Some(serde_json::Value::Object(extra)) = extra_args.as_mut() { // These fields are already represented by token_ids and media. @@ -110,13 +163,15 @@ pub(crate) fn build_generate_request( ignore_eos: stop_conditions.ignore_eos.unwrap_or(false), }), response: Some(pb::ResponseOptions { - prompt_token_ids: prompt_logprobs.is_some() || (has_media && mode.is_prefill()), + prompt_token_ids: prompt_logprobs.is_some() || (has_raw_media && mode.is_prefill()), prompt_logprobs: prompt_logprobs.is_some(), prompt_candidates: prompt_logprobs.map(top_n_candidates).transpose()?, output_text: Some(true), output_token_ids: true, output_logprobs: output_logprobs.is_some(), output_candidates: output_logprobs.map(top_n_candidates).transpose()?, + skip_special_tokens: request.output_options.skip_special_tokens, + routed_experts_prompt_start: vllm_tito.routed_experts_prompt_start, }), kv: Some(kv), truncate_prompt_tokens: 0, @@ -126,6 +181,257 @@ pub(crate) fn build_generate_request( }) } +fn validate_and_remove_vllm_tito( + extra_args: &mut Option, + canonical_cache_salt: Option<&str>, + canonical_priority: i32, + canonical_request_id: &str, +) -> Result { + let Some(serde_json::Value::Object(extra)) = extra_args.as_mut() else { + return Ok(VllmTitoProjection::default()); + }; + let Some(envelope) = extra.remove("vllm_tito") else { + return Ok(VllmTitoProjection::default()); + }; + let serde_json::Value::Object(envelope) = envelope else { + return Err(client::invalid_argument( + "extra_args.vllm_tito must be a JSON object", + )); + }; + + for key in envelope.keys() { + if !matches!( + key.as_str(), + "request_id" + | "sampling_params" + | "model" + | "stream" + | "stream_options" + | "cache_salt" + | "priority" + | "kv_transfer_params" + | "features" + ) { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.{key} is not supported by vLLM gRPC" + ))); + } + } + + if envelope + .get("request_id") + .and_then(serde_json::Value::as_str) + .is_none_or(|request_id| request_id != canonical_request_id) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.request_id must match the canonical request ID", + )); + } + + let features = envelope + .get("features") + .filter(|features| !features.is_null()) + .map(|features| { + serde_json::from_value::(features.clone()).map_err(|error| { + client::invalid_argument(format!( + "extra_args.vllm_tito.features is invalid: {error}" + )) + }) + }) + .transpose()?; + if envelope + .get("model") + .is_some_and(|model| !model.is_null() && !model.is_string()) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.model must be a string", + )); + } + if envelope + .get("stream") + .is_some_and(|stream| stream != &serde_json::Value::Bool(false)) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.stream must be false", + )); + } + if envelope + .get("stream_options") + .is_some_and(|options| !options.is_null()) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.stream_options is not supported by vLLM gRPC", + )); + } + + let sampling = envelope.get("sampling_params").ok_or_else(|| { + client::invalid_argument("extra_args.vllm_tito.sampling_params is required") + })?; + let serde_json::Value::Object(sampling) = sampling else { + return Err(client::invalid_argument( + "extra_args.vllm_tito.sampling_params must be a JSON object", + )); + }; + for key in sampling.keys() { + if !matches!( + key.as_str(), + "temperature" + | "top_p" + | "top_k" + | "min_p" + | "seed" + | "max_tokens" + | "min_tokens" + | "presence_penalty" + | "frequency_penalty" + | "repetition_penalty" + | "stop" + | "stop_token_ids" + | "ignore_eos" + | "logprobs" + | "prompt_logprobs" + | "cache_salt" + | "skip_reading_prefix_cache" + | "skip_special_tokens" + | "include_stop_str_in_output" + | "return_token_ids" + | "routed_experts_prompt_start" + ) { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.sampling_params.{key} is not supported by vLLM gRPC" + ))); + } + } + if sampling + .get("skip_special_tokens") + .is_some_and(|value| !value.is_boolean()) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.sampling_params.skip_special_tokens must be a boolean", + )); + } + let routed_experts_prompt_start = sampling + .get("routed_experts_prompt_start") + .filter(|value| !value.is_null()) + .map(|value| { + value + .as_u64() + .and_then(|value| u32::try_from(value).ok()) + .ok_or_else(|| { + client::invalid_argument( + "extra_args.vllm_tito.sampling_params.routed_experts_prompt_start must be an unsigned 32-bit integer", + ) + }) + }) + .transpose()?; + if sampling + .get("return_token_ids") + .is_some_and(|value| value != &serde_json::Value::Bool(true)) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.sampling_params.return_token_ids must be true", + )); + } + validate_compat_cache_salt( + sampling.get("cache_salt"), + canonical_cache_salt, + "sampling_params.cache_salt", + )?; + validate_compat_cache_salt( + envelope.get("cache_salt"), + canonical_cache_salt, + "cache_salt", + )?; + + if let Some(skip_reading_prefix_cache) = sampling.get("skip_reading_prefix_cache") { + if !skip_reading_prefix_cache.is_boolean() { + return Err(client::invalid_argument( + "extra_args.vllm_tito.sampling_params.skip_reading_prefix_cache must be a boolean", + )); + } + insert_compatible_projection( + extra, + "skip_reading_prefix_cache", + skip_reading_prefix_cache.clone(), + "sampling_params.skip_reading_prefix_cache", + )?; + } + if let Some(kv_transfer_params) = envelope.get("kv_transfer_params") + && !kv_transfer_params.is_null() + { + if !kv_transfer_params.is_object() { + return Err(client::invalid_argument( + "extra_args.vllm_tito.kv_transfer_params must be a JSON object", + )); + } + insert_compatible_projection( + extra, + "kv_transfer_params", + kv_transfer_params.clone(), + "kv_transfer_params", + )?; + } + + let priority = envelope + .get("priority") + .and_then(serde_json::Value::as_i64) + .and_then(|priority| i32::try_from(priority).ok()) + .ok_or_else(|| { + client::invalid_argument( + "extra_args.vllm_tito.priority must be a signed 32-bit integer", + ) + })?; + if priority.saturating_neg() != canonical_priority { + return Err(client::invalid_argument( + "extra_args.vllm_tito.priority does not match the canonical Dynamo routing priority", + )); + } + Ok(VllmTitoProjection { + priority: Some(priority), + features, + routed_experts_prompt_start, + }) +} + +fn insert_compatible_projection( + extra: &mut serde_json::Map, + key: &str, + value: serde_json::Value, + envelope_path: &str, +) -> Result<(), DynamoError> { + if let Some(existing) = extra.get(key) { + if existing != &value { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.{envelope_path} conflicts with extra_args.{key}" + ))); + } + } else { + extra.insert(key.to_string(), value); + } + Ok(()) +} + +fn validate_compat_cache_salt( + value: Option<&serde_json::Value>, + canonical: Option<&str>, + path: &str, +) -> Result<(), DynamoError> { + let Some(value) = value else { + return Ok(()); + }; + let Some(value) = value.as_str() else { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.{path} must be a string" + ))); + }; + if Some(value) != canonical { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.{path} must match the canonical cache_salt" + ))); + } + Ok(()) +} + pub(crate) fn data_parallel_rank( request: &PreprocessedRequest, mode: DisaggregationMode, @@ -358,7 +664,118 @@ fn build_media( Ok(media) } +fn build_preprocessed_media(features: VllmTitoFeatures) -> Result, DynamoError> { + if features.mm_hashes.is_empty() { + return Err(client::invalid_argument( + "extra_args.vllm_tito.features.mm_hashes must contain at least one modality", + )); + } + if features + .mm_hashes + .keys() + .ne(features.mm_placeholders.keys()) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.features hashes and placeholders must contain the same modalities", + )); + } + if let Some(kwargs_data) = &features.kwargs_data + && features.mm_hashes.keys().ne(kwargs_data.keys()) + { + return Err(client::invalid_argument( + "extra_args.vllm_tito.features kwargs_data must contain the same modalities as mm_hashes", + )); + } + + let mut media = Vec::new(); + for (modality, hashes) in features.mm_hashes { + let placeholders = features + .mm_placeholders + .get(&modality) + .expect("validated modality keys"); + if hashes.len() != placeholders.len() { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features.{modality} hashes and placeholders must have equal lengths" + ))); + } + let kwargs = features + .kwargs_data + .as_ref() + .map(|kwargs_data| &kwargs_data[&modality]); + if kwargs.is_some_and(|kwargs| kwargs.len() != hashes.len()) { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features.kwargs_data.{modality} must align with hashes and placeholders" + ))); + } + let modality_code = match modality.as_str() { + "image" => pb::Modality::Image, + "video" => pb::Modality::Video, + "audio" => pb::Modality::Audio, + _ => { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features modality `{modality}` is not supported" + ))); + } + }; + + for (index, (identifier, placeholder)) in hashes.into_iter().zip(placeholders).enumerate() { + if identifier.is_empty() { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features.mm_hashes.{modality}[{index}] must be non-empty" + ))); + } + if placeholder.length == 0 { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features.mm_placeholders.{modality}[{index}].length must be positive" + ))); + } + if placeholder + .is_embed + .as_ref() + .is_some_and(|mask| mask.len() as u64 != placeholder.length) + { + return Err(client::invalid_argument(format!( + "extra_args.vllm_tito.features.mm_placeholders.{modality}[{index}].is_embed must match length" + ))); + } + let encoded_kwargs = kwargs + .and_then(|kwargs| kwargs.get(index)) + .and_then(|kwargs| kwargs.as_deref()); + let decoded_kwargs = encoded_kwargs + .map(|kwargs| { + BASE64_STANDARD.decode(kwargs).map_err(|error| { + client::invalid_argument(format!( + "extra_args.vllm_tito.features.kwargs_data.{modality}[{index}] is not valid base64: {error}" + )) + }) + }) + .transpose()?; + media.push(pb::MediaItem { + modality: modality_code as i32, + source: Some(pb::media_item::Source::Features( + pb::PreprocessedMediaFeatures { + kwargs: decoded_kwargs, + identifier: identifier.clone(), + offset: placeholder.offset, + length: placeholder.length, + mm_hash: Some(identifier), + is_embed: placeholder.is_embed.clone().unwrap_or_default(), + }, + )), + mime_type: String::new(), + uuid: String::new(), + }); + } + } + Ok(media) +} + fn top_n_candidates(count: u32) -> Result { + if count == u32::MAX { + return Ok(pb::CandidateTokens { + select: Some(pb::candidate_tokens::Select::All(true)), + }); + } i32::try_from(count).map_err(|_| { client::invalid_argument(format!( "vLLM logprobs request must fit in i32; got {count}" @@ -605,11 +1022,6 @@ fn validate_request( "max_thinking_tokens is not supported by vLLM gRPC v0.25.1", )); } - if request.output_options.skip_special_tokens == Some(false) { - return Err(client::invalid_argument( - "skip_special_tokens=false is not supported by vLLM gRPC v0.25.1", - )); - } let sampling = &request.sampling_options; if sampling.n.unwrap_or(1) != 1 { return Err(client::invalid_argument("n must be 1")); @@ -701,6 +1113,11 @@ impl ResponseState { } else { None }; + let routed_experts = output + .routed_experts + .as_ref() + .map(routed_experts_to_json) + .transpose()?; let pb::SequenceOutput { text, num_tokens, @@ -730,6 +1147,11 @@ impl ResponseState { } let Some(finish) = finish_info else { + if routed_experts.is_some() { + return Err(client::protocol_error( + "routed_experts are only valid on a terminal sequence output", + )); + } return if self.is_prefill || num_tokens == 0 { Ok(None) } else { @@ -795,6 +1217,15 @@ impl ResponseState { ); } self.attach_prompt_data(&mut mapped); + if let Some(routed_experts) = routed_experts { + let engine_data = mapped + .engine_data + .get_or_insert_with(|| serde_json::Value::Object(serde_json::Map::new())); + let engine_data = engine_data.as_object_mut().ok_or_else(|| { + client::protocol_error("terminal engine_data is not a JSON object") + })?; + engine_data.insert("routed_experts".to_string(), routed_experts); + } Ok(Some(mapped)) } @@ -848,6 +1279,44 @@ impl ResponseState { } } +fn routed_experts_to_json(routed: &pb::RoutedExperts) -> Result { + if routed.shape.len() != 3 { + return Err(client::protocol_error(format!( + "routed_experts shape must have rank 3, got {:?}", + routed.shape + ))); + } + let item_size = match routed.dtype.as_str() { + "uint8" => 1usize, + "uint16" => 2usize, + dtype => { + return Err(client::protocol_error(format!( + "routed_experts dtype must be uint8 or uint16, got {dtype:?}" + ))); + } + }; + let expected = routed + .shape + .iter() + .try_fold(1usize, |count, dimension| { + count.checked_mul(*dimension as usize) + }) + .and_then(|count| count.checked_mul(item_size)) + .ok_or_else(|| client::protocol_error("routed_experts byte length overflow"))?; + if routed.data.len() != expected { + return Err(client::protocol_error(format!( + "routed_experts byte length mismatch: expected {expected}, got {}", + routed.data.len() + ))); + } + Ok(serde_json::json!({ + "data": BASE64_STANDARD.encode(&routed.data), + "shape": routed.shape, + "start": routed.start, + "dtype": routed.dtype, + })) +} + fn prompt_logprobs_to_json(prompt: pb::PromptInfo) -> serde_json::Value { let count = prompt.num_prompt_tokens as usize; let mut positions = Vec::with_capacity(count); diff --git a/lib/sidecar/vllm/src/engine.rs b/lib/sidecar/vllm/src/engine.rs index 5c7936f50f0e..a6e5d5e75640 100644 --- a/lib/sidecar/vllm/src/engine.rs +++ b/lib/sidecar/vllm/src/engine.rs @@ -101,6 +101,12 @@ impl VllmSidecarEngine { } let endpoint = GrpcEndpoint::parse(&args.vllm_endpoint, "--vllm-endpoint")?; + let vllm_http_url = args + .vllm_http_endpoint + .as_deref() + .map(|value| GrpcEndpoint::parse(value, "--vllm-http-endpoint")) + .transpose()? + .map(|endpoint| endpoint.to_string()); let transport = args.sidecar.grpc.config(); let bootstrap_deadline = client::startup_deadline(transport.startup_deadline)?; eprintln!( @@ -109,6 +115,12 @@ impl VllmSidecarEngine { ); let model = bootstrap_discover(&endpoint, transport, bootstrap_deadline)?; let mode = args.sidecar.common.disaggregation_mode; + let rl_metadata = args + .sidecar + .common + .enable_rl + .then(|| model.rl_worker_metadata(vllm_http_url)) + .transpose()?; let engine = Self::new(endpoint, model.clone(), mode, transport); let config = WorkerConfig { namespace: args.sidecar.common.namespace, @@ -135,6 +147,7 @@ impl VllmSidecarEngine { disaggregation_mode: mode, route_to_encoder: false, enable_rl: args.sidecar.common.enable_rl, + rl_metadata, ..Default::default() }; Ok((engine, config)) diff --git a/lib/sidecar/vllm/src/model.rs b/lib/sidecar/vllm/src/model.rs index 313ee3da146e..088a25b7808f 100644 --- a/lib/sidecar/vllm/src/model.rs +++ b/lib/sidecar/vllm/src/model.rs @@ -1,12 +1,14 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -use dynamo_backend_common::{DynamoError, EngineConfig, LlmRegistration}; +use dynamo_backend_common::{DynamoError, EngineConfig, LlmRegistration, RlWorkerMetadata}; use crate::client; use crate::proto as pb; const SUPPORTED_API_VERSION: &str = "vllm"; +const VLLM_INFERENCE_V1_GENERATE_CAPABILITY: &str = "vllm_inference_v1_generate"; +const VLLM_EXACT_MM_ROUTING_CAPABILITY: &str = "vllm_exact_mm_routing"; #[derive(Clone, Debug, Eq, PartialEq)] struct ModelIdentity { @@ -89,6 +91,12 @@ impl DiscoveredModel { "data-parallel size changed between bootstrap and startup: expected {expected_dp_size}, observed {observed_dp_size}" ))); } + if self.server.parallelism != observed.server.parallelism { + return Err(client::protocol_error(format!( + "parallelism changed between bootstrap and startup: expected {:?}, observed {:?}", + self.server.parallelism, observed.server.parallelism + ))); + } if self.server.rl_capabilities != observed.server.rl_capabilities { return Err(client::protocol_error(format!( "RL capabilities changed between bootstrap and startup: expected {:?}, observed {:?}", @@ -102,13 +110,49 @@ impl DiscoveredModel { self.server.rl_capabilities.as_ref() } + pub(crate) fn rl_worker_metadata( + &self, + admin_base_url: Option, + ) -> Result { + let parallelism = self.server.parallelism.as_ref().ok_or_else(|| { + client::protocol_error("RL discovery requires vLLM parallelism metadata") + })?; + let world_size = parallelism + .tensor_parallel_size + .checked_mul(parallelism.pipeline_parallel_size) + .and_then(|size| size.checked_mul(parallelism.data_parallel_size)) + .filter(|size| *size > 0) + .ok_or_else(|| client::protocol_error("vLLM reports an invalid RL world size"))?; + let backend = self + .server + .rl_capabilities + .as_ref() + .and_then(|capabilities| { + capabilities + .weight_transfer_enabled + .then(|| capabilities.weight_transfer_backend.clone()) + }); + RlWorkerMetadata::new(world_size, backend, admin_base_url) + .map_err(|error| client::protocol_error(error.to_string())) + } + pub(crate) fn engine_config(&self) -> EngineConfig { let parallelism = self.server.parallelism.as_ref(); EngineConfig { model: self.source.clone(), served_model_name: Some(self.served_name.clone()), model_aliases: self.identity.aliases.clone(), - runtime_data: Default::default(), + runtime_data: [ + ( + VLLM_INFERENCE_V1_GENERATE_CAPABILITY.to_string(), + serde_json::Value::Bool(true), + ), + ( + VLLM_EXACT_MM_ROUTING_CAPABILITY.to_string(), + serde_json::Value::Bool(self.supports_multimodal), + ), + ] + .into(), llm: Some(LlmRegistration { context_length: nonzero(self.server.max_model_len), kv_cache_block_size: nonzero(self.server.kv_block_size), diff --git a/lib/sidecar/vllm/src/tests.rs b/lib/sidecar/vllm/src/tests.rs index e29576249fcb..70a6a5cd29b9 100644 --- a/lib/sidecar/vllm/src/tests.rs +++ b/lib/sidecar/vllm/src/tests.rs @@ -11,7 +11,7 @@ use std::sync::atomic::{AtomicBool, Ordering}; use dynamo_backend_common::engine::RoutingHints; use dynamo_backend_common::{ DisaggregationMode, FinishReason, GenerateContext, LLMEngine, MultimodalData, OutputOptions, - PrefillResult, PreprocessedRequest, SamplingOptions, StopConditions, + PrefillResult, PreprocessedRequest, RlWorkerMetadata, SamplingOptions, StopConditions, }; use dynamo_sidecar_common::{GrpcEndpoint, GrpcTransportConfig}; use futures::{Stream, StreamExt}; @@ -459,6 +459,49 @@ fn server_info() -> pb::ServerInfo { } } +#[test] +fn discovered_model_reports_rl_worker_metadata() { + let model = DiscoveredModel::from_proto(model_info(), server_info()).expect("valid discovery"); + assert_eq!( + model + .rl_worker_metadata(Some("http://worker:8120".to_string())) + .expect("valid RL metadata"), + RlWorkerMetadata::new( + 4, + Some("nccl".to_string()), + Some("http://worker:8120".to_string()), + ) + .expect("valid expected metadata") + ); +} + +#[test] +fn discovered_model_advertises_token_native_generate() { + let model = DiscoveredModel::from_proto(model_info(), server_info()).expect("valid discovery"); + assert_eq!( + model + .engine_config() + .runtime_data + .get("vllm_inference_v1_generate"), + Some(&serde_json::Value::Bool(true)) + ); +} + +#[test] +fn startup_compatibility_rejects_parallelism_change() { + let expected = + DiscoveredModel::from_proto(model_info(), server_info()).expect("valid discovery"); + let mut changed_server = server_info(); + changed_server + .parallelism + .as_mut() + .expect("parallelism") + .tensor_parallel_size = 4; + let changed = + DiscoveredModel::from_proto(model_info(), changed_server).expect("valid changed discovery"); + assert!(expected.ensure_startup_compatible(&changed).is_err()); +} + fn sequence_response( terminal: bool, logprobs: bool, @@ -489,6 +532,7 @@ fn sequence_response( kv_transfer_params, ec_transfer_params: None, }), + routed_experts: None, }), } } @@ -585,6 +629,30 @@ fn zero_output_logprobs_omits_top_logprobs() { assert!(mapped.top_logprobs.is_none()); } +#[test] +fn all_logprobs_use_the_proto_all_selector() { + let mut request = request(); + request.output_options.logprobs = Some(u32::MAX); + request.output_options.prompt_logprobs = Some(u32::MAX); + + let wire = build_generate_request( + request, + "all-logprobs".to_string(), + DisaggregationMode::Aggregated, + ) + .expect("build request"); + let response = wire.response.expect("response options"); + + assert!(matches!( + response.output_candidates.and_then(|value| value.select), + Some(pb::candidate_tokens::Select::All(true)) + )); + assert!(matches!( + response.prompt_candidates.and_then(|value| value.select), + Some(pb::candidate_tokens::Select::All(true)) + )); +} + #[test] fn oversized_logprob_counts_are_rejected() { let oversized = i32::MAX as u32 + 1; @@ -610,6 +678,229 @@ fn oversized_logprob_counts_are_rejected() { assert!(prompt_error.to_string().contains("must fit in i32")); } +#[test] +fn token_native_compatibility_envelope_is_accepted_when_fields_are_projected() { + let mut request = request(); + request.output_options.skip_special_tokens = Some(false); + request.routing.as_mut().unwrap().priority = Some(-7); + request.extra_args = Some(json!({ + "vllm_tito": { + "request_id": "request-1", + "sampling_params": { + "temperature": 0.2, + "top_p": 0.9, + "top_k": 4, + "min_p": 0.1, + "seed": 123, + "max_tokens": 1, + "min_tokens": 1, + "presence_penalty": 0.3, + "frequency_penalty": 0.4, + "repetition_penalty": 1.1, + "stop_token_ids": [2], + "ignore_eos": true, + "logprobs": 1, + "prompt_logprobs": 1, + "cache_salt": "cache-salt", + "skip_reading_prefix_cache": true, + "skip_special_tokens": false, + "return_token_ids": true, + "routed_experts_prompt_start": 2 + }, + "model": "served-model", + "stream": false, + "cache_salt": "cache-salt", + "priority": 7, + "kv_transfer_params": { + "connector_data": {"values": [1, true, null]} + } + } + })); + + let wire = build_generate_request( + request, + "request-1".to_string(), + DisaggregationMode::Aggregated, + ) + .expect("projected compatibility envelope"); + + assert_eq!(wire.temperature, Some(0.2)); + assert_eq!(wire.stopping.as_ref().unwrap().stop_token_ids, [2]); + assert_eq!( + wire.response.as_ref().unwrap().skip_special_tokens, + Some(false) + ); + assert!(wire.response.as_ref().unwrap().output_logprobs); + assert_eq!( + wire.kv.as_ref().unwrap().cache_salt, + "dynamo-cache-salt:cache-salt" + ); + assert!(wire.kv.as_ref().unwrap().bypass_prefix_cache); + assert!(wire.kv.as_ref().unwrap().kv_transfer_params.is_some()); + assert_eq!(wire.priority, 7); + assert_eq!( + wire.response.as_ref().unwrap().routed_experts_prompt_start, + Some(2) + ); +} + +#[test] +fn terminal_routed_experts_are_mapped_to_structured_engine_data() { + let request = request(); + let mut state = ResponseState::new(&request, DisaggregationMode::Aggregated); + let mut response = sequence_response(true, true, None); + response.outputs.as_mut().unwrap().routed_experts = Some(pb::RoutedExperts { + data: vec![1, 2, 3, 4], + shape: vec![2, 1, 2], + dtype: "uint8".to_string(), + start: 3, + }); + + let output = state + .convert(response) + .expect("valid routed experts") + .expect("terminal output"); + assert_eq!( + output + .engine_data + .as_ref() + .and_then(|data| data.get("routed_experts")), + Some(&json!({ + "data": "AQIDBA==", + "shape": [2, 1, 2], + "start": 3, + "dtype": "uint8" + })) + ); +} + +#[test] +fn malformed_terminal_routed_experts_are_rejected() { + let request = request(); + let mut state = ResponseState::new(&request, DisaggregationMode::Aggregated); + let mut response = sequence_response(true, true, None); + response.outputs.as_mut().unwrap().routed_experts = Some(pb::RoutedExperts { + data: vec![1], + shape: vec![2, 1, 2], + dtype: "uint8".to_string(), + start: 0, + }); + + let error = state + .convert(response) + .expect_err("byte-length mismatch must fail"); + assert!(error.to_string().contains("byte length mismatch")); +} + +#[test] +fn token_native_compatibility_envelope_rejects_disabled_token_ids() { + let mut request = request(); + request.extra_args = Some(json!({ + "vllm_tito": { + "request_id": "request-1", + "sampling_params": { + "max_tokens": 1, + "return_token_ids": false + }, + "stream": false, + "priority": 0 + } + })); + + let error = build_generate_request( + request, + "request-1".to_string(), + DisaggregationMode::Aggregated, + ) + .expect_err("the token-native sidecar always returns token IDs"); + + assert!( + error + .to_string() + .contains("sampling_params.return_token_ids must be true") + ); +} + +#[test] +fn token_native_compatibility_envelope_rejects_conflicting_cache_salt() { + let mut request = request(); + request.extra_args = Some(json!({ + "vllm_tito": { + "request_id": "request-1", + "sampling_params": { + "max_tokens": 1, + "cache_salt": "different-policy-version" + }, + "cache_salt": "cache-salt", + "stream": false, + "priority": 0 + } + })); + + let error = build_generate_request( + request, + "request-1".to_string(), + DisaggregationMode::Aggregated, + ) + .expect_err("duplicate cache salts must match the canonical routing salt"); + + assert!( + error + .to_string() + .contains("sampling_params.cache_salt must match the canonical cache_salt") + ); +} + +#[test] +fn token_native_compatibility_envelope_rejects_unprojected_fields() { + let mut request = request(); + request.extra_args = Some(json!({ + "vllm_tito": { + "request_id": "request-1", + "sampling_params": { + "max_tokens": 1, + "future_sampling_field": true + }, + "stream": false, + "priority": 0 + } + })); + + let error = build_generate_request( + request, + "request-1".to_string(), + DisaggregationMode::Aggregated, + ) + .expect_err("unprojected compatibility fields must fail closed"); + + assert!( + error + .to_string() + .contains("extra_args.vllm_tito.sampling_params.future_sampling_field") + ); +} + +#[test] +fn skip_special_tokens_false_is_forwarded_to_vllm() { + let mut request = request(); + request.output_options.skip_special_tokens = Some(false); + + let mapped = build_generate_request( + request, + "skip-special-tokens".to_string(), + DisaggregationMode::Aggregated, + ) + .expect("skip_special_tokens=false should be supported"); + + assert_eq!( + mapped + .response + .expect("response options") + .skip_special_tokens, + Some(false) + ); +} + struct FakeServer { endpoint: String, service: FakeVllm, @@ -726,6 +1017,58 @@ fn decode_request() -> PreprocessedRequest { request } +fn request_with_preprocessed_image(hash: &str, encoded_kwargs: &str) -> PreprocessedRequest { + let mut request = request(); + request.token_ids = vec![10, 99, 99, 20]; + request.extra_args = Some(json!({ + "vllm_tito": { + "request_id": "feature-request", + "sampling_params": {"cache_salt": "cache-salt"}, + "stream": false, + "cache_salt": "cache-salt", + "priority": 0, + "features": { + "mm_hashes": {"image": [hash]}, + "mm_placeholders": {"image": [{"offset": 1, "length": 2}]}, + "kwargs_data": {"image": [encoded_kwargs]} + } + } + })); + request +} + +#[test] +fn preprocessed_images_with_same_layout_keep_distinct_execution_identity() { + let first = build_generate_request( + request_with_preprocessed_image("image-red", "cmVk"), + "feature-request".to_string(), + DisaggregationMode::Aggregated, + ) + .expect("first preprocessed image"); + let second = build_generate_request( + request_with_preprocessed_image("image-blue", "Ymx1ZQ=="), + "feature-request".to_string(), + DisaggregationMode::Aggregated, + ) + .expect("second preprocessed image"); + + assert_eq!(first.prompt, second.prompt); + let first_feature = match first.media[0].source.as_ref() { + Some(pb::media_item::Source::Features(feature)) => feature, + other => panic!("expected preprocessed features, got {other:?}"), + }; + let second_feature = match second.media[0].source.as_ref() { + Some(pb::media_item::Source::Features(feature)) => feature, + other => panic!("expected preprocessed features, got {other:?}"), + }; + assert_eq!((first_feature.offset, first_feature.length), (1, 2)); + assert_eq!((second_feature.offset, second_feature.length), (1, 2)); + assert_eq!(first_feature.identifier, "image-red"); + assert_eq!(second_feature.identifier, "image-blue"); + assert_eq!(first_feature.kwargs.as_deref(), Some(&b"red"[..])); + assert_eq!(second_feature.kwargs.as_deref(), Some(&b"blue"[..])); +} + fn engine( endpoint: &str, mode: DisaggregationMode, @@ -761,6 +1104,9 @@ async fn engine_from_args( "dynamo-vllm-sidecar".to_string(), "--vllm-endpoint".to_string(), endpoint.to_string(), + "--vllm-http-endpoint".to_string(), + "http://worker:8120".to_string(), + "--enable-rl".to_string(), "--grpc-connections".to_string(), "2".to_string(), "--grpc-startup-deadline-secs".to_string(), @@ -835,6 +1181,17 @@ async fn aggregated_generation_converts_request_stream_and_usage() { assert_eq!(worker.served_model_name.as_deref(), Some("served-model")); assert!(worker.reasoning_parser.is_none()); assert!(worker.tool_call_parser.is_none()); + assert_eq!( + worker.rl_metadata, + Some( + RlWorkerMetadata::new( + 4, + Some("nccl".to_string()), + Some("http://worker:8120".to_string()), + ) + .expect("valid RL metadata") + ) + ); let config = engine.start(0).await.expect("start"); assert_eq!(config.model, "model-source"); assert_eq!(config.served_model_name.as_deref(), Some("served-model"));