Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions rust/src/server/src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,7 @@ fn is_request_validation_error(error: &vllm_text::Error) -> bool {
vllm_text::Error::PromptTooLong { .. }
| vllm_text::Error::EmptyPromptTokenIds { .. }
| vllm_text::Error::Logprobs(_)
| vllm_text::Error::OutOfVocab(_)
// An empty tokenized prompt detected later, at request prepare
// time, surfaces through the transparent Llm wrapper.
| vllm_text::Error::Llm(vllm_llm::Error::EmptyPromptTokenIds { .. })
Expand Down Expand Up @@ -172,6 +173,17 @@ mod tests {
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
}

#[test]
fn out_of_vocab_validation_maps_to_invalid_request() {
let error = vllm_text::Error::OutOfVocab(vllm_text::OutOfVocabError {
parameter: "logprob_token_ids",
token_ids: vec![1000],
vocab_size: 1000,
});
let api_error = text_submit_error("failed to submit completion request", error);
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
}

#[test]
fn other_submit_errors_stay_internal() {
let error = vllm_text::Error::Tokenizer("backend exploded".to_string());
Expand Down
7 changes: 0 additions & 7 deletions rust/src/server/src/routes/openai/chat_completions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,13 +53,6 @@ pub async fn chat_completions(
let request_context = resolve_request_context(&headers, body.request_id.as_deref());
let lora_resolution = state.resolve_model_with_loras(Some(&body.model)).await;

if let Err(err) = validate::validate_token_id_ranges(
&body,
state.tokenizer_vocab_size(),
state.model_vocab_size(),
) {
return err.into_response();
}
let prepared = match prepare_chat_request(body, &lora_resolution, request_context) {
Ok(prepared) => prepared,
Err(error) => return error.into_response(),
Expand Down
44 changes: 1 addition & 43 deletions rust/src/server/src/routes/openai/chat_completions/validate.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
use super::types::ChatCompletionRequest;
use crate::error::{ApiError, bail_invalid_request};
use crate::routes::openai::utils::token_ids::{validate_allowed_token_ids, validate_logit_bias};
use crate::routes::openai::utils::types::{ChatMessage, Tool, ToolChoice, ToolChoiceValue};

/// Enforce the minimal compatibility contract for the Rust OpenAI server.
Expand Down Expand Up @@ -154,29 +153,14 @@ fn validate_function_tools(tools: &[Tool], param: &'static str) -> Result<(), Ap
Ok(())
}

/// Reject out-of-vocab token ids, mirroring the Python input processor:
/// `allowed_token_ids` against the tokenizer vocab, `logit_bias` keys against the
/// model vocab (skipped when the model size is unknown).
pub(super) fn validate_token_id_ranges(
request: &ChatCompletionRequest,
tokenizer_vocab_size: usize,
model_vocab_size: Option<usize>,
) -> Result<(), ApiError> {
validate_allowed_token_ids(request.allowed_token_ids.as_deref(), tokenizer_vocab_size)?;
validate_logit_bias(
request.logit_bias.as_ref(),
model_vocab_size.unwrap_or(usize::MAX),
)
}

#[cfg(test)]
mod tests {
use std::collections::HashMap;

use serde_json::json;
use vllm_chat::ReasoningEffort;

use super::{validate_request_compat, validate_token_id_ranges};
use super::validate_request_compat;
use crate::routes::openai::chat_completions::types::ChatCompletionRequest;
use crate::routes::openai::utils::structured_outputs::ResponseFormat;
use crate::routes::openai::utils::types::{
Expand All @@ -188,32 +172,6 @@ mod tests {
names.iter().map(|s| s.to_string()).collect()
}

#[test]
fn validate_token_id_ranges_rejects_oob_and_accepts_in_vocab() {
// allowed_token_ids are bounded by the tokenizer vocab
let mut request = base_request();
request.allowed_token_ids = Some(vec![5, 1_000_000]);
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
// logit_bias is bounded by the larger model vocab: an id between the two
// vocabs is valid and must not be rejected (the parity regression we fix)
let mut request = base_request();
request.logit_bias = Some(HashMap::from([("150".to_string(), 1.0)]));
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
// logit_bias beyond the model vocab -> reject
let mut request = base_request();
request.logit_bias = Some(HashMap::from([("1000000".to_string(), 1.0)]));
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
// all in-vocab -> accept
let mut request = base_request();
request.allowed_token_ids = Some(vec![5, 50]);
request.logit_bias = Some(HashMap::from([("50".to_string(), 1.0)]));
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
// unknown sizes -> skip
let mut request = base_request();
request.allowed_token_ids = Some(vec![1_000_000]);
assert!(validate_token_id_ranges(&request, usize::MAX, None).is_ok());
}

fn base_request() -> ChatCompletionRequest {
ChatCompletionRequest {
model: "Qwen/Qwen1.5-0.5B-Chat".to_string(),
Expand Down
7 changes: 0 additions & 7 deletions rust/src/server/src/routes/openai/completions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -47,13 +47,6 @@ pub async fn completions(
let request_context = resolve_request_context(&headers, body.request_id.as_deref());
let lora_resolution = state.resolve_model_with_loras(Some(&body.model)).await;

if let Err(err) = validate::validate_token_id_ranges(
&body,
state.tokenizer_vocab_size(),
state.model_vocab_size(),
) {
return err.into_response();
}
let prepared = match prepare_completion_request(body, &lora_resolution, request_context) {
Ok(prepared) => prepared,
Err(error) => return error.into_response(),
Expand Down
55 changes: 1 addition & 54 deletions rust/src/server/src/routes/openai/completions/validate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,6 @@ use vllm_text::Prompt;

use super::types::CompletionRequest;
use crate::error::{ApiError, bail_invalid_request};
use crate::routes::openai::utils::token_ids::{
validate_allowed_token_ids, validate_logit_bias, validate_prompt_token_ids,
};

/// Enforce the minimal compatibility contract for the Rust OpenAI server.
pub(super) fn validate_request_compat(
Expand Down Expand Up @@ -98,63 +95,13 @@ pub(super) fn validate_request_compat(
Ok(())
}

/// Reject out-of-vocab token ids, mirroring the Python input processor. A token-id
/// prompt may reference ids the engine embeds beyond either vocab alone (Qwen3
/// extra LM tokens, multimodal placeholders), so it is bounded by the union of the
/// tokenizer and model vocabularies; `allowed_token_ids` by the tokenizer vocab;
/// `logit_bias` keys by the model vocab (skipped when the model size is unknown).
pub(super) fn validate_token_id_ranges(
request: &CompletionRequest,
tokenizer_vocab_size: usize,
model_vocab_size: Option<usize>,
) -> Result<(), ApiError> {
let prompt_bound = tokenizer_vocab_size.max(model_vocab_size.unwrap_or(0));
validate_prompt_token_ids(&request.prompt, prompt_bound)?;
validate_allowed_token_ids(request.allowed_token_ids.as_deref(), tokenizer_vocab_size)?;
validate_logit_bias(
request.logit_bias.as_ref(),
model_vocab_size.unwrap_or(usize::MAX),
)
}

#[cfg(test)]
mod tests {
use serde_json::json;
use vllm_text::Prompt;

use super::{validate_request_compat, validate_token_id_ranges};
use super::validate_request_compat;
use crate::routes::openai::completions::types::CompletionRequest;

#[test]
fn validate_token_id_ranges_rejects_oob_prompt_and_params() {
// a token-id prompt below both vocabs is accepted (the engine can embed it)
let mut request = base_request();
request.prompt = Prompt::TokenIds(vec![5, 150]);
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
// an id at or above the union of the two vocabs is rejected
let mut request = base_request();
request.prompt = Prompt::TokenIds(vec![5, 200]);
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
// an id beyond the model vocab but within the (larger) tokenizer vocab is
// accepted: the engine embeds added/placeholder ids above the model vocab,
// matching the Python input processor's max(tokenizer, model) bound
let mut request = base_request();
request.prompt = Prompt::TokenIds(vec![150]);
assert!(validate_token_id_ranges(&request, 200, Some(100)).is_ok());
// falls back to the tokenizer vocab when the model size is unknown
let mut request = base_request();
request.prompt = Prompt::TokenIds(vec![150]);
assert!(validate_token_id_ranges(&request, 100, None).is_err());
// allowed_token_ids are bounded by the tokenizer vocab -> reject
let mut request = base_request();
request.allowed_token_ids = Some(vec![150]);
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
// unknown sizes -> skip
let mut request = base_request();
request.prompt = Prompt::TokenIds(vec![1_000_000]);
assert!(validate_token_id_ranges(&request, usize::MAX, None).is_ok());
}

fn base_request() -> CompletionRequest {
serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
Expand Down
1 change: 0 additions & 1 deletion rust/src/server/src/routes/openai/utils/mod.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
pub mod logprobs;
pub mod structured_outputs;
pub mod token_ids;
pub mod types;
pub mod usage;
pub mod validated_json;
55 changes: 0 additions & 55 deletions rust/src/server/src/routes/openai/utils/token_ids.rs

This file was deleted.

10 changes: 0 additions & 10 deletions rust/src/server/src/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -114,16 +114,6 @@ impl AppState {
&self.served_model_names
}

/// Tokenizer vocabulary size.
pub fn tokenizer_vocab_size(&self) -> usize {
self.chat.tokenizer_vocab_size()
}

/// Model vocabulary size, else `None`.
pub fn model_vocab_size(&self) -> Option<usize> {
self.chat.model_vocab_size()
}

/// Return base served model names plus dynamically loaded LoRA adapter
/// names.
pub async fn served_model_names_with_loras(&self) -> Vec<String> {
Expand Down
18 changes: 15 additions & 3 deletions rust/src/text/src/backend/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -29,10 +29,12 @@ pub struct SamplingLimits {
///
/// `-1` means allowing requests up to the model vocabulary size.
pub max_logprobs: i32,
/// Model vocabulary size from the model config.

/// Model vocabulary size from the model config, used to bound
/// `logit_bias` keys when available.
pub model_vocab_size: Option<usize>,
/// Tokenizer vocabulary size, used as a fallback when the model config does
/// not expose a vocabulary size.
/// Tokenizer vocabulary size, used to bound `allowed_token_ids` and
/// token-ID prompts.
pub tokenizer_vocab_size: usize,
}

Expand All @@ -48,6 +50,16 @@ impl SamplingLimits {
pub fn logprobs_vocab_size(&self) -> usize {
self.model_vocab_size.unwrap_or(self.tokenizer_vocab_size)
}

/// Return the vocabulary size used to validate generated stop token IDs.
pub fn stop_token_vocab_size(&self) -> usize {
self.model_vocab_size.unwrap_or(self.tokenizer_vocab_size)
}

/// Return the union bound used to validate token-ID prompts.
pub fn prompt_token_vocab_size(&self) -> usize {
self.tokenizer_vocab_size.max(self.model_vocab_size.unwrap_or(0))
}
}

/// Minimal text-processing backend needed by `vllm-text`.
Expand Down
3 changes: 3 additions & 0 deletions rust/src/text/src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ use vllm_engine_core_client::Error as EngineCoreError;
use vllm_llm::Error as LlmError;

pub use crate::lower::logprobs::LogprobsError;
pub use crate::lower::token_ids::OutOfVocabError;

#[derive(Debug, Error)]
pub enum Error {
Expand All @@ -17,6 +18,8 @@ pub enum Error {
PromptTooLong { max_model_len: u32, prompt_len: u32 },
#[error(transparent)]
Logprobs(#[from] LogprobsError),
#[error(transparent)]
OutOfVocab(#[from] OutOfVocabError),
#[error("text request stream `{request_id}` closed before terminal output")]
StreamClosedBeforeTerminalOutput { request_id: String },
#[error(transparent)]
Expand Down
2 changes: 1 addition & 1 deletion rust/src/text/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
use std::mem::take;

pub use backend::{DynTextBackend, SamplingHints, SamplingLimits, TextBackend};
pub use error::{Error, LogprobsError, Result};
pub use error::{Error, LogprobsError, OutOfVocabError, Result};
use futures::Stream;
pub use lower::{
PreparedTextRequest, lower_sampling_params, lower_text_request, resolve_max_tokens,
Expand Down
Loading
Loading