mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-14 17:58:11 +00:00
feat(frontend): session id plumbing into requests (#48048)
Signed-off-by: Karen Chung <[email protected]>
This commit is contained in:
@@ -45,6 +45,8 @@ message GenerateRequest {
|
||||
uint32 truncate_prompt_tokens = 11;
|
||||
|
||||
int32 priority = 12;
|
||||
|
||||
optional string session_id = 13;
|
||||
}
|
||||
|
||||
message RandomSampling {
|
||||
|
||||
@@ -168,6 +168,7 @@ impl ChatRequestProcessor {
|
||||
cache_salt: request.cache_salt,
|
||||
add_special_tokens: request.add_special_tokens,
|
||||
data_parallel_rank: request.data_parallel_rank,
|
||||
session_id: request.session_id,
|
||||
reasoning_parser_kwargs,
|
||||
lora_request: request.lora_request,
|
||||
arrival_time: Some(arrival_time),
|
||||
|
||||
@@ -491,6 +491,9 @@ pub struct ChatRequest {
|
||||
/// Override data parallel rank.
|
||||
#[serde(default)]
|
||||
pub data_parallel_rank: Option<u32>,
|
||||
/// Stable session identity shared by related requests.
|
||||
#[serde(default)]
|
||||
pub session_id: Option<String>,
|
||||
/// LoRA adapter selected for this request.
|
||||
#[serde(default)]
|
||||
pub lora_request: Option<LoraRequest>,
|
||||
@@ -514,6 +517,7 @@ impl ChatRequest {
|
||||
cache_salt: None,
|
||||
add_special_tokens: false,
|
||||
data_parallel_rank: None,
|
||||
session_id: None,
|
||||
lora_request: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -124,6 +124,9 @@ pub struct EngineCoreRequest {
|
||||
/// standard `request_finished` hook.
|
||||
#[serde(default)]
|
||||
pub abort_immediately: bool,
|
||||
/// Stable session identity shared by related requests.
|
||||
#[serde(default)]
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
impl EngineCoreRequest {
|
||||
@@ -175,6 +178,7 @@ mod tests {
|
||||
}),
|
||||
arrival_time: 1234.5,
|
||||
client_index: 7,
|
||||
session_id: Some("session-1".to_string()),
|
||||
..EngineCoreRequest::default()
|
||||
};
|
||||
|
||||
@@ -185,12 +189,13 @@ mod tests {
|
||||
other => panic!("expected array, got {other:?}"),
|
||||
};
|
||||
|
||||
assert_eq!(array.len(), 20);
|
||||
assert_eq!(array.len(), 21);
|
||||
assert_eq!(array[0], Value::from("req-1"));
|
||||
assert_eq!(array[2], Value::Nil);
|
||||
assert_eq!(array[4], Value::Nil);
|
||||
assert_eq!(array[10], Value::Nil);
|
||||
assert_eq!(array[11], Value::from(7));
|
||||
assert_eq!(array[20], Value::from("session-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -162,6 +162,7 @@ fn sample_request_with_id(request_id: &str) -> EngineCoreRequest {
|
||||
..EngineCoreSamplingParams::for_test()
|
||||
}),
|
||||
arrival_time: 42.5,
|
||||
session_id: Some("session-1".to_string()),
|
||||
..EngineCoreRequest::default()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -75,6 +75,7 @@ class EngineCoreRequest(
|
||||
reasoning_ended: bool | None = None
|
||||
reasoning_parser_kwargs: dict[str, object] | None = None
|
||||
abort_immediately: bool = False
|
||||
session_id: str | None = None
|
||||
|
||||
|
||||
class EngineCoreOutput(
|
||||
@@ -137,6 +138,7 @@ request = EngineCoreRequest(
|
||||
pooling_params=None,
|
||||
arrival_time=42.5,
|
||||
client_index=0,
|
||||
session_id="session-1",
|
||||
)
|
||||
|
||||
# All defaults -> empty map. Regression guard for the sparse-map decode.
|
||||
|
||||
@@ -59,6 +59,7 @@ fn build_request(request_id: String, max_tokens: u32) -> GenerateRequest {
|
||||
trace_headers: None,
|
||||
priority: 0,
|
||||
data_parallel_rank: None,
|
||||
session_id: None,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: None,
|
||||
}
|
||||
|
||||
@@ -46,6 +46,8 @@ pub struct GenerateRequest {
|
||||
pub priority: i32,
|
||||
/// Optional data-parallel rank override for routing this request.
|
||||
pub data_parallel_rank: Option<u32>,
|
||||
/// Stable session identity shared by related requests.
|
||||
pub session_id: Option<String>,
|
||||
/// Optional reasoning-parser kwargs forwarded to engine-side structured
|
||||
/// output logic.
|
||||
pub reasoning_parser_kwargs: Option<ReasoningParserKwargs>,
|
||||
@@ -76,6 +78,7 @@ impl GenerateRequest {
|
||||
trace_headers,
|
||||
priority,
|
||||
data_parallel_rank,
|
||||
session_id,
|
||||
reasoning_parser_kwargs,
|
||||
lora_request,
|
||||
} = self;
|
||||
@@ -105,6 +108,7 @@ impl GenerateRequest {
|
||||
priority,
|
||||
trace_headers,
|
||||
resumable: false,
|
||||
session_id,
|
||||
external_req_id: Some(external_request_id),
|
||||
// Rust parser doesn't expose this information, leave it unset and let the
|
||||
// reasoning logic in engine-sided structured output manager handle it.
|
||||
@@ -150,6 +154,7 @@ mod tests {
|
||||
)])),
|
||||
priority: 3,
|
||||
data_parallel_rank: Some(2),
|
||||
session_id: Some("session-1".to_string()),
|
||||
reasoning_parser_kwargs: Some(ReasoningParserKwargs {
|
||||
chat_template_kwargs: [(
|
||||
"chat_template_kwargs".to_string(),
|
||||
@@ -177,6 +182,7 @@ mod tests {
|
||||
assert_eq!(request.arrival_time, 42.5);
|
||||
assert_eq!(request.cache_salt.as_deref(), Some("salt"));
|
||||
assert_eq!(request.data_parallel_rank, Some(2));
|
||||
assert_eq!(request.session_id.as_deref(), Some("session-1"));
|
||||
assert_eq!(
|
||||
request.trace_headers,
|
||||
Some(BTreeMap::from([(
|
||||
|
||||
@@ -188,6 +188,7 @@ fn sample_generate_request(request_id: &str, max_tokens: u32) -> GenerateRequest
|
||||
trace_headers: None,
|
||||
priority: 0,
|
||||
data_parallel_rank: None,
|
||||
session_id: None,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: None,
|
||||
}
|
||||
|
||||
@@ -53,6 +53,7 @@ pub fn to_text_request(
|
||||
} else {
|
||||
req.request_id
|
||||
};
|
||||
let session_id = req.session_id.filter(|s| !s.is_empty());
|
||||
|
||||
let sampling = req.sampling.as_ref();
|
||||
let decoding = req.decoding.as_ref();
|
||||
@@ -100,6 +101,7 @@ pub fn to_text_request(
|
||||
cache_salt: kv.map(|k| &k.cache_salt).filter(|s| !s.is_empty()).cloned(),
|
||||
add_special_tokens: true,
|
||||
data_parallel_rank: None,
|
||||
session_id,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: None,
|
||||
arrival_time: None,
|
||||
|
||||
@@ -75,6 +75,7 @@ pub(super) fn prepare_generate_request(
|
||||
cache_salt: request.cache_salt,
|
||||
add_special_tokens: false,
|
||||
data_parallel_rank: ctx.data_parallel_rank,
|
||||
session_id: ctx.session_id,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: lora_resolution.lora_request.clone(),
|
||||
arrival_time: None,
|
||||
|
||||
@@ -18,6 +18,7 @@ use crate::routes::openai::utils::types::{
|
||||
};
|
||||
use crate::utils::{
|
||||
ResolvedRequestContext, convert_logit_bias, merge_ec_transfer_params, merge_kv_transfer_params,
|
||||
resolve_session_id,
|
||||
};
|
||||
|
||||
/// Lowered chat request plus the public response metadata carried by every SSE
|
||||
@@ -123,6 +124,11 @@ pub(super) fn prepare_chat_request(
|
||||
request.response_format.as_ref(),
|
||||
&request.structured_outputs,
|
||||
)?;
|
||||
let session_id = resolve_session_id(
|
||||
&ctx,
|
||||
request.session_id.as_deref(),
|
||||
request.vllm_xargs.as_ref(),
|
||||
);
|
||||
|
||||
let chat_request = ChatRequest {
|
||||
request_id: request_id.clone(),
|
||||
@@ -177,6 +183,7 @@ pub(super) fn prepare_chat_request(
|
||||
cache_salt: request.cache_salt,
|
||||
add_special_tokens: request.add_special_tokens,
|
||||
data_parallel_rank: ctx.data_parallel_rank,
|
||||
session_id,
|
||||
lora_request: lora_resolution.lora_request.clone(),
|
||||
};
|
||||
|
||||
@@ -1241,6 +1248,49 @@ mod tests {
|
||||
assert_eq!(prepared.chat_request.data_parallel_rank, Some(7));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_chat_request_threads_header_session_id() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("X-Session-ID", "header-session".parse().unwrap());
|
||||
let prepared = prepare_chat_request(
|
||||
base_request(),
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
request_context(&headers, None),
|
||||
)
|
||||
.expect("request is valid");
|
||||
assert_eq!(
|
||||
prepared.chat_request.session_id.as_deref(),
|
||||
Some("header-session")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_chat_request_ignores_correlation_header() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("X-Correlation-ID", "correlation-header".parse().unwrap());
|
||||
let prepared = prepare_chat_request(
|
||||
base_request(),
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
request_context(&headers, None),
|
||||
)
|
||||
.expect("request is valid");
|
||||
assert_eq!(prepared.chat_request.session_id, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_chat_request_ignores_empty_session_header_without_fallback() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("X-Session-ID", "".parse().unwrap());
|
||||
headers.insert("X-Correlation-ID", "correlation-header".parse().unwrap());
|
||||
let prepared = prepare_chat_request(
|
||||
base_request(),
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
request_context(&headers, None),
|
||||
)
|
||||
.expect("request is valid");
|
||||
assert_eq!(prepared.chat_request.session_id, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_chat_request_leaves_data_parallel_rank_none_when_absent() {
|
||||
let prepared = prepare_chat_request(
|
||||
|
||||
@@ -222,6 +222,9 @@ pub struct ChatCompletionRequest {
|
||||
/// External request ID used for response correlation.
|
||||
pub request_id: Option<String>,
|
||||
|
||||
/// Stable session identity shared by related requests.
|
||||
pub session_id: Option<String>,
|
||||
|
||||
/// Tokens represented as strings of the form 'token_id:{token_id}' in
|
||||
/// logprobs
|
||||
pub return_tokens_as_token_ids: Option<bool>,
|
||||
@@ -301,6 +304,7 @@ impl Default for ChatCompletionRequest {
|
||||
structured_outputs: None,
|
||||
priority: None,
|
||||
request_id: None,
|
||||
session_id: None,
|
||||
return_tokens_as_token_ids: None,
|
||||
return_token_ids: None,
|
||||
cache_salt: None,
|
||||
|
||||
@@ -12,6 +12,7 @@ use crate::routes::openai::completions::validate;
|
||||
use crate::routes::openai::utils::structured_outputs::convert_from_response_format_value;
|
||||
use crate::utils::{
|
||||
ResolvedRequestContext, convert_logit_bias, merge_ec_transfer_params, merge_kv_transfer_params,
|
||||
resolve_session_id,
|
||||
};
|
||||
|
||||
/// Lowered completion request plus the public response metadata carried by
|
||||
@@ -103,6 +104,11 @@ pub(super) fn prepare_completion_request(
|
||||
|
||||
let structured_outputs =
|
||||
convert_from_response_format_value(&request.response_format, &request.structured_outputs)?;
|
||||
let session_id = resolve_session_id(
|
||||
&ctx,
|
||||
request.session_id.as_deref(),
|
||||
request.vllm_xargs.as_ref(),
|
||||
);
|
||||
|
||||
let text_request = TextRequest {
|
||||
request_id: request_id.clone(),
|
||||
@@ -147,6 +153,7 @@ pub(super) fn prepare_completion_request(
|
||||
cache_salt: request.cache_salt,
|
||||
add_special_tokens: request.add_special_tokens,
|
||||
data_parallel_rank: ctx.data_parallel_rank,
|
||||
session_id,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: lora_resolution.lora_request.clone(),
|
||||
arrival_time: None,
|
||||
@@ -586,6 +593,76 @@ mod tests {
|
||||
assert_eq!(prepared.text_request.data_parallel_rank, Some(3));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_threads_body_session_id() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"prompt": "hello",
|
||||
"stream": false,
|
||||
"session_id": "body-session",
|
||||
"vllm_xargs": {"session_id": "xargs-session"},
|
||||
}))
|
||||
.expect("parse request");
|
||||
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("X-Session-ID", "header-session".parse().unwrap());
|
||||
let prepared = prepare_completion_request(
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
request_context(&headers, None),
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(
|
||||
prepared.text_request.session_id.as_deref(),
|
||||
Some("body-session")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_uses_vllm_xargs_session_id_fallback() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"prompt": "hello",
|
||||
"stream": false,
|
||||
"vllm_xargs": {"session_id": "xargs-session"},
|
||||
}))
|
||||
.expect("parse request");
|
||||
|
||||
let prepared = prepare_completion_request(
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(
|
||||
prepared.text_request.session_id.as_deref(),
|
||||
Some("xargs-session")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_ignores_empty_and_non_string_session_id_values() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"prompt": "hello",
|
||||
"stream": false,
|
||||
"session_id": "",
|
||||
"vllm_xargs": {"session_id": 7},
|
||||
}))
|
||||
.expect("parse request");
|
||||
|
||||
let prepared = prepare_completion_request(
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(prepared.text_request.session_id, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_leaves_data_parallel_rank_none_when_absent() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
|
||||
@@ -164,6 +164,9 @@ pub struct CompletionRequest {
|
||||
/// External request ID used for response correlation.
|
||||
pub request_id: Option<String>,
|
||||
|
||||
/// Stable session identity shared by related requests.
|
||||
pub session_id: Option<String>,
|
||||
|
||||
/// Tokens represented as strings of the form 'token_id:{token_id}' in
|
||||
/// logprobs
|
||||
pub return_tokens_as_token_ids: Option<bool>,
|
||||
|
||||
@@ -95,6 +95,7 @@ impl TokenizeChatRequest {
|
||||
cache_salt: None,
|
||||
add_special_tokens: self.add_special_tokens,
|
||||
data_parallel_rank: None,
|
||||
session_id: None,
|
||||
lora_request: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@ use crate::error::ApiError;
|
||||
pub struct ResolvedRequestContext {
|
||||
pub request_id: String,
|
||||
pub data_parallel_rank: Option<u32>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
/// Return the current Unix timestamp in seconds for OpenAI response objects.
|
||||
@@ -66,6 +67,24 @@ pub fn merge_ec_transfer_params(
|
||||
xargs
|
||||
}
|
||||
|
||||
pub fn resolve_session_id(
|
||||
ctx: &ResolvedRequestContext,
|
||||
request_session_id: Option<&str>,
|
||||
xargs: Option<&HashMap<String, Value>>,
|
||||
) -> Option<String> {
|
||||
request_session_id
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_owned)
|
||||
.or_else(|| ctx.session_id.clone())
|
||||
.or_else(|| {
|
||||
xargs
|
||||
.and_then(|map| map.get("session_id"))
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert OpenAI-style `logit_bias` with string token-ID keys into the
|
||||
/// internal `HashMap<u32, f32>` representation, validating that every key
|
||||
/// parses as a `u32`.
|
||||
@@ -91,8 +110,8 @@ pub fn convert_logit_bias(
|
||||
.transpose()
|
||||
}
|
||||
|
||||
/// Extract common request metadata from HTTP headers: the external request ID
|
||||
/// and the optional data-parallel rank used for engine routing.
|
||||
/// Extract common request metadata from HTTP headers: the external request ID,
|
||||
/// session ID, and the optional data-parallel rank used for engine routing.
|
||||
pub fn resolve_request_context(
|
||||
headers: &HeaderMap,
|
||||
request_id: Option<&str>,
|
||||
@@ -106,10 +125,16 @@ pub fn resolve_request_context(
|
||||
// Extract request id from header.
|
||||
let request_id_header = headers.get("X-Request-Id").and_then(|value| value.to_str().ok());
|
||||
let request_id = resolve_base_request_id(request_id_header, request_id);
|
||||
let session_id = headers
|
||||
.get("X-Session-ID")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_owned);
|
||||
|
||||
ResolvedRequestContext {
|
||||
request_id,
|
||||
data_parallel_rank,
|
||||
session_id,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -59,6 +59,7 @@ pub fn lower_text_request(
|
||||
cache_salt: request.cache_salt.clone(),
|
||||
priority: request.priority,
|
||||
data_parallel_rank: request.data_parallel_rank,
|
||||
session_id: request.session_id.clone(),
|
||||
reasoning_parser_kwargs: request.reasoning_parser_kwargs.clone(),
|
||||
lora_request: request.lora_request.clone(),
|
||||
arrival_time: request.arrival_time,
|
||||
|
||||
@@ -183,6 +183,9 @@ pub struct TextRequest {
|
||||
/// Override data parallel rank.
|
||||
#[serde(default)]
|
||||
pub data_parallel_rank: Option<u32>,
|
||||
/// Stable session identity shared by related requests.
|
||||
#[serde(default)]
|
||||
pub session_id: Option<String>,
|
||||
/// Optional reasoning-parser kwargs forwarded to engine-side structured
|
||||
/// output logic.
|
||||
#[serde(default)]
|
||||
@@ -212,6 +215,7 @@ impl TextRequest {
|
||||
cache_salt: None,
|
||||
add_special_tokens: false,
|
||||
data_parallel_rank: None,
|
||||
session_id: None,
|
||||
reasoning_parser_kwargs: None,
|
||||
lora_request: None,
|
||||
arrival_time: None,
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import pytest
|
||||
from starlette.requests import Request
|
||||
|
||||
from vllm.entrypoints.generate.base.serving import GenerateBaseServing
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.completion.protocol import CompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
|
||||
def _raw_request(headers: dict[str, str]) -> Request:
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/chat/completions",
|
||||
"headers": [
|
||||
(key.lower().encode("latin-1"), value.encode("latin-1"))
|
||||
for key, value in headers.items()
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"openai_request",
|
||||
[
|
||||
ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
),
|
||||
CompletionRequest(model="test-model", prompt="hi"),
|
||||
ResponsesRequest(model="test-model", input="hi"),
|
||||
],
|
||||
)
|
||||
def test_get_session_id_accepts_body_field(openai_request):
|
||||
openai_request.session_id = "body-session"
|
||||
|
||||
session_id = GenerateBaseServing._get_session_id(
|
||||
openai_request,
|
||||
_raw_request({"X-Session-ID": "header-session"}),
|
||||
)
|
||||
|
||||
assert session_id == "body-session"
|
||||
|
||||
|
||||
def test_get_session_id_accepts_session_header():
|
||||
request = CompletionRequest(model="test-model", prompt="hi")
|
||||
|
||||
session_id = GenerateBaseServing._get_session_id(
|
||||
request,
|
||||
_raw_request({"X-Session-ID": "header-session"}),
|
||||
)
|
||||
|
||||
assert session_id == "header-session"
|
||||
|
||||
|
||||
def test_get_session_id_ignores_correlation_header():
|
||||
request = CompletionRequest(
|
||||
model="test-model",
|
||||
prompt="hi",
|
||||
vllm_xargs={"session_id": "xargs-session"},
|
||||
)
|
||||
|
||||
session_id = GenerateBaseServing._get_session_id(
|
||||
request,
|
||||
_raw_request({"X-Correlation-ID": "correlation-session"}),
|
||||
)
|
||||
|
||||
assert session_id == "xargs-session"
|
||||
|
||||
|
||||
def test_get_session_id_keeps_vllm_xargs_as_compatibility_fallback():
|
||||
request = CompletionRequest(
|
||||
model="test-model",
|
||||
prompt="hi",
|
||||
vllm_xargs={"session_id": "xargs-session"},
|
||||
)
|
||||
|
||||
session_id = GenerateBaseServing._get_session_id(request, None)
|
||||
|
||||
assert session_id == "xargs-session"
|
||||
|
||||
|
||||
def test_get_session_id_ignores_empty_and_non_string_values():
|
||||
request = CompletionRequest(
|
||||
model="test-model",
|
||||
prompt="hi",
|
||||
session_id="",
|
||||
vllm_xargs={"session_id": 7},
|
||||
)
|
||||
|
||||
session_id = GenerateBaseServing._get_session_id(request, None)
|
||||
|
||||
assert session_id is None
|
||||
@@ -203,6 +203,32 @@ async def test_serve_tokens_skips_mm_cache_for_remote_engine_execution():
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serve_tokens_threads_session_id_header_to_engine():
|
||||
engine = _mock_engine()
|
||||
|
||||
async def mock_generate(*args, **kwargs):
|
||||
yield _make_request_output(
|
||||
"req-1", token_ids=[10], finish_reason="stop", finished=True
|
||||
)
|
||||
|
||||
engine.generate = MagicMock(side_effect=mock_generate)
|
||||
serving = _build_serving_tokens(engine)
|
||||
|
||||
request = GenerateRequest(
|
||||
token_ids=[1, 2, 3],
|
||||
sampling_params=SamplingParams(max_tokens=1),
|
||||
model=MODEL_NAME,
|
||||
stream=False,
|
||||
)
|
||||
raw_request = MagicMock()
|
||||
raw_request.headers = {"X-Session-ID": "header-session"}
|
||||
|
||||
await serving.serve_tokens(request, raw_request)
|
||||
|
||||
assert engine.generate.call_args.kwargs["session_id"] == "header-session"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_basic():
|
||||
"""Streaming returns SSE chunks with correct token_ids and ends with [DONE]."""
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from copy import copy
|
||||
|
||||
from vllm import SamplingParams
|
||||
from vllm.outputs import CompletionOutput
|
||||
from vllm.sampling_params import RequestOutputKind
|
||||
@@ -68,6 +70,20 @@ def test_parent_request_to_output_final_only() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_parallel_sampling_child_requests_preserve_session_id() -> None:
|
||||
request = make_request(SamplingParams(n=2))
|
||||
request.session_id = "session-1"
|
||||
parent_request = ParentRequest(request)
|
||||
|
||||
for idx in range(parent_request.n):
|
||||
request_id, child_params = parent_request.get_child_info(idx)
|
||||
child_request = request if idx == parent_request.n - 1 else copy(request)
|
||||
child_request.request_id = request_id
|
||||
child_request.sampling_params = child_params
|
||||
|
||||
assert child_request.session_id == "session-1"
|
||||
|
||||
|
||||
def make_request(sampling_params: SamplingParams) -> EngineCoreRequest:
|
||||
return EngineCoreRequest(
|
||||
request_id="parent_id",
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from vllm.v1.request import RequestStatus
|
||||
|
||||
from vllm import SamplingParams
|
||||
from vllm.v1.engine import EngineCoreRequest
|
||||
from vllm.v1.request import Request, RequestStatus
|
||||
|
||||
|
||||
def test_request_status_fmt_str():
|
||||
@@ -18,3 +21,22 @@ def test_request_status_fmt_str():
|
||||
assert f"{RequestStatus.FINISHED_LENGTH_CAPPED}" == "FINISHED_LENGTH_CAPPED"
|
||||
assert f"{RequestStatus.FINISHED_ABORTED}" == "FINISHED_ABORTED"
|
||||
assert f"{RequestStatus.FINISHED_IGNORED}" == "FINISHED_IGNORED"
|
||||
|
||||
|
||||
def test_request_copies_session_id_from_engine_core_request():
|
||||
engine_request = EngineCoreRequest(
|
||||
request_id="request-1",
|
||||
prompt_token_ids=[1, 2, 3],
|
||||
mm_features=None,
|
||||
sampling_params=SamplingParams(max_tokens=1),
|
||||
pooling_params=None,
|
||||
arrival_time=0.0,
|
||||
lora_request=None,
|
||||
cache_salt=None,
|
||||
data_parallel_rank=None,
|
||||
session_id="session-1",
|
||||
)
|
||||
|
||||
request = Request.from_engine_core_request(engine_request, block_hasher=None)
|
||||
|
||||
assert request.session_id == "session-1"
|
||||
|
||||
@@ -78,6 +78,7 @@ class EngineClient(ABC):
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
priority: int = 0,
|
||||
data_parallel_rank: int | None = None,
|
||||
session_id: str | None = None,
|
||||
reasoning_ended: bool | None = None,
|
||||
reasoning_parser_kwargs: dict[str, Any] | None = None,
|
||||
) -> AsyncGenerator[RequestOutput, None]:
|
||||
|
||||
@@ -41,6 +41,7 @@ logger = init_logger(__name__)
|
||||
|
||||
RequestT = TypeVar("RequestT", bound=AnyRequest)
|
||||
_T = TypeVar("_T")
|
||||
SESSION_ID_HEADER = "X-Session-ID"
|
||||
|
||||
|
||||
def build_per_request_timing_metrics(
|
||||
@@ -215,6 +216,29 @@ class GenerateBaseServing(BaseServing, BeamSearchOnlineMixin):
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_headers(raw_request: Request | None) -> str | None:
|
||||
if raw_request is None:
|
||||
return None
|
||||
if value := raw_request.headers.get(SESSION_ID_HEADER):
|
||||
return value
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id(
|
||||
request: ChatCompletionRequest | CompletionRequest | ResponsesRequest,
|
||||
raw_request: Request | None,
|
||||
) -> str | None:
|
||||
if request.session_id:
|
||||
return request.session_id
|
||||
if value := GenerateBaseServing._get_session_id_from_headers(raw_request):
|
||||
return value
|
||||
if request.vllm_xargs:
|
||||
session_id = request.vllm_xargs.get("session_id")
|
||||
if isinstance(session_id, str) and session_id:
|
||||
return session_id
|
||||
return None
|
||||
|
||||
async def _with_kv_transfer_rejection_cleanup(
|
||||
self,
|
||||
awaitable: Awaitable[_T],
|
||||
|
||||
@@ -32,6 +32,7 @@ class BeamSearchOnlineMixin(ABC):
|
||||
params: BeamSearchParams,
|
||||
lora_request: LoRARequest | None = None,
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> AsyncGenerator[RequestOutput, None]:
|
||||
beam_width = params.beam_width
|
||||
max_tokens = params.max_tokens
|
||||
@@ -90,6 +91,7 @@ class BeamSearchOnlineMixin(ABC):
|
||||
request_id_item,
|
||||
lora_request=lora_request_item,
|
||||
trace_headers=trace_headers,
|
||||
session_id=session_id,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
@@ -174,6 +174,7 @@ class OpenAIServingChatBatch(OpenAIServingChat):
|
||||
if raw_request is None
|
||||
else await self._get_trace_headers(raw_request.headers)
|
||||
)
|
||||
session_id = self._get_session_id(single_request, raw_request)
|
||||
generators.append(
|
||||
self.engine_client.generate(
|
||||
engine_prompt,
|
||||
@@ -183,6 +184,7 @@ class OpenAIServingChatBatch(OpenAIServingChat):
|
||||
trace_headers=trace_headers,
|
||||
priority=request.priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
reasoning_ended=None,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -394,6 +394,14 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
"through out the inference process and return in response."
|
||||
),
|
||||
)
|
||||
session_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Stable session identity shared by related requests. Unlike "
|
||||
"request_id, this value is expected to remain stable across "
|
||||
"multiple requests in the same conversation or agent session."
|
||||
),
|
||||
)
|
||||
|
||||
return_tokens_as_token_ids: bool | None = Field(
|
||||
default=None,
|
||||
|
||||
@@ -318,6 +318,7 @@ class OpenAIServingChat(GenerateBaseServing):
|
||||
if raw_request is None
|
||||
else await self._get_trace_headers(raw_request.headers)
|
||||
)
|
||||
session_id = self._get_session_id(request, raw_request)
|
||||
|
||||
if isinstance(sampling_params, BeamSearchParams):
|
||||
generator = self.beam_search(
|
||||
@@ -326,6 +327,7 @@ class OpenAIServingChat(GenerateBaseServing):
|
||||
params=sampling_params,
|
||||
lora_request=lora_request,
|
||||
trace_headers=trace_headers,
|
||||
session_id=session_id,
|
||||
)
|
||||
else:
|
||||
if not request.include_reasoning:
|
||||
@@ -348,6 +350,7 @@ class OpenAIServingChat(GenerateBaseServing):
|
||||
trace_headers=trace_headers,
|
||||
priority=request.priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
reasoning_ended=reasoning_ended,
|
||||
reasoning_parser_kwargs={
|
||||
"chat_template_kwargs": chat_template_kwargs,
|
||||
|
||||
@@ -147,6 +147,14 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
"through out the inference process and return in response."
|
||||
),
|
||||
)
|
||||
session_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Stable session identity shared by related requests. Unlike "
|
||||
"request_id, this value is expected to remain stable across "
|
||||
"multiple requests in the same conversation or agent session."
|
||||
),
|
||||
)
|
||||
|
||||
return_tokens_as_token_ids: bool | None = Field(
|
||||
default=None,
|
||||
|
||||
@@ -194,6 +194,7 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
if raw_request is None
|
||||
else await self._get_trace_headers(raw_request.headers)
|
||||
)
|
||||
session_id = self._get_session_id(request, raw_request)
|
||||
|
||||
if isinstance(sampling_params, BeamSearchParams):
|
||||
generator = self.beam_search(
|
||||
@@ -202,6 +203,7 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
params=sampling_params,
|
||||
lora_request=lora_request,
|
||||
trace_headers=trace_headers,
|
||||
session_id=session_id,
|
||||
)
|
||||
else:
|
||||
generator = self.engine_client.generate(
|
||||
@@ -212,6 +214,7 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
trace_headers=trace_headers,
|
||||
priority=request.priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
)
|
||||
|
||||
generators.append(generator)
|
||||
|
||||
@@ -220,6 +220,14 @@ class ResponsesRequest(OpenAIBaseModel):
|
||||
"through out the inference process and return in response."
|
||||
),
|
||||
)
|
||||
session_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Stable session identity shared by related requests. Unlike "
|
||||
"request_id, this value is expected to remain stable across "
|
||||
"multiple requests in the same conversation or agent session."
|
||||
),
|
||||
)
|
||||
media_io_kwargs: dict[str, dict[str, Any]] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
||||
@@ -451,6 +451,7 @@ class OpenAIServingResponses(GenerateBaseServing):
|
||||
if raw_request is None
|
||||
else await self._get_trace_headers(raw_request.headers)
|
||||
)
|
||||
session_id = self._get_session_id(request, raw_request)
|
||||
|
||||
chat_template_kwargs = self._effective_chat_template_kwargs(request)
|
||||
response_parser = self._make_response_parser(
|
||||
@@ -516,6 +517,7 @@ class OpenAIServingResponses(GenerateBaseServing):
|
||||
lora_request=lora_request,
|
||||
priority=request.priority,
|
||||
trace_headers=trace_headers,
|
||||
session_id=session_id,
|
||||
reasoning_parser_kwargs=reasoning_parser_kwargs
|
||||
if self.parser and self.parser.reasoning_parser_cls is not None
|
||||
else None,
|
||||
@@ -669,6 +671,7 @@ class OpenAIServingResponses(GenerateBaseServing):
|
||||
lora_request: LoRARequest | None = None,
|
||||
priority: int = 0,
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
session_id: str | None = None,
|
||||
reasoning_parser_kwargs: dict[str, Any] | None = None,
|
||||
):
|
||||
max_model_len = self.model_config.max_model_len
|
||||
@@ -693,6 +696,7 @@ class OpenAIServingResponses(GenerateBaseServing):
|
||||
lora_request=lora_request,
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
session_id=session_id,
|
||||
reasoning_parser_kwargs=reasoning_parser_kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -215,6 +215,7 @@ class ServingTokens(GenerateBaseServing):
|
||||
|
||||
# Extract data_parallel_rank from header (router can inject it)
|
||||
data_parallel_rank = self._get_data_parallel_rank(raw_request)
|
||||
session_id = self._get_session_id_from_headers(raw_request)
|
||||
|
||||
result_generator = self.engine_client.generate(
|
||||
engine_input,
|
||||
@@ -224,6 +225,7 @@ class ServingTokens(GenerateBaseServing):
|
||||
trace_headers=trace_headers,
|
||||
priority=request.priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
)
|
||||
|
||||
assert result_generator is not None
|
||||
|
||||
@@ -145,6 +145,8 @@ class EngineCoreRequest(
|
||||
# KV-transfer request is rejected on the D node before engine admission.
|
||||
abort_immediately: bool = False
|
||||
|
||||
session_id: str | None = None
|
||||
|
||||
@property
|
||||
def params(self) -> SamplingParams | PoolingParams:
|
||||
"""Return the processed params (sampling or pooling)."""
|
||||
|
||||
@@ -294,6 +294,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
priority: int = 0,
|
||||
data_parallel_rank: int | None = None,
|
||||
session_id: str | None = None,
|
||||
prompt_text: str | None = None,
|
||||
reasoning_ended: bool | None = None,
|
||||
reasoning_parser_kwargs: dict[str, Any] | None = None,
|
||||
@@ -331,6 +332,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers,
|
||||
priority,
|
||||
data_parallel_rank,
|
||||
session_id,
|
||||
)
|
||||
|
||||
# Convert Input --> Request.
|
||||
@@ -362,6 +364,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
)
|
||||
else:
|
||||
# Raw prompts require tokenization and possibly multimodal
|
||||
@@ -377,6 +380,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
)
|
||||
prompt_text, _, _ = extract_prompt_components(self.model_config, prompt)
|
||||
|
||||
@@ -445,6 +449,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
priority: int = 0,
|
||||
data_parallel_rank: int | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> RequestOutputCollector:
|
||||
self._validate_streaming_input_sampling_params(sampling_params)
|
||||
|
||||
@@ -456,6 +461,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
)
|
||||
|
||||
if not sampling_params.skip_clone:
|
||||
@@ -556,6 +562,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
priority: int = 0,
|
||||
data_parallel_rank: int | None = None,
|
||||
session_id: str | None = None,
|
||||
reasoning_ended: bool | None = None,
|
||||
reasoning_parser_kwargs: dict[str, Any] | None = None,
|
||||
) -> AsyncGenerator[RequestOutput, None]:
|
||||
@@ -585,6 +592,7 @@ class AsyncLLM(EngineClient):
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
session_id=session_id,
|
||||
prompt_text=prompt_text,
|
||||
reasoning_ended=reasoning_ended,
|
||||
reasoning_parser_kwargs=reasoning_parser_kwargs,
|
||||
|
||||
@@ -261,6 +261,7 @@ class InputProcessor:
|
||||
priority: int = 0,
|
||||
data_parallel_rank: int | None = None,
|
||||
resumable: bool = False,
|
||||
session_id: str | None = None,
|
||||
) -> EngineCoreRequest:
|
||||
self._validate_params(params, supported_tasks)
|
||||
self._validate_lora(lora_request)
|
||||
@@ -391,6 +392,7 @@ class InputProcessor:
|
||||
data_parallel_rank=data_parallel_rank,
|
||||
trace_headers=trace_headers,
|
||||
resumable=resumable,
|
||||
session_id=session_id,
|
||||
)
|
||||
|
||||
def _validate_prompt_len(
|
||||
|
||||
@@ -225,6 +225,7 @@ class LLMEngine:
|
||||
tokenization_kwargs: dict[str, Any] | None = None,
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
priority: int = 0,
|
||||
session_id: str | None = None,
|
||||
prompt_text: str | None = None,
|
||||
) -> str:
|
||||
# Validate the request_id type.
|
||||
@@ -257,6 +258,7 @@ class LLMEngine:
|
||||
tokenization_kwargs=tokenization_kwargs,
|
||||
trace_headers=trace_headers,
|
||||
priority=priority,
|
||||
session_id=session_id,
|
||||
)
|
||||
prompt_text, _, _ = extract_prompt_components(self.model_config, prompt)
|
||||
|
||||
|
||||
@@ -74,6 +74,7 @@ class Request:
|
||||
trace_headers: Mapping[str, str] | None = None,
|
||||
block_hasher: Callable[["Request"], list["BlockHash"]] | None = None,
|
||||
resumable: bool = False,
|
||||
session_id: str | None = None,
|
||||
reasoning_ended: bool | None = None,
|
||||
reasoning_parser_kwargs: dict[str, Any] | None = None,
|
||||
abort_immediately: bool = False,
|
||||
@@ -183,6 +184,7 @@ class Request:
|
||||
self.all_token_ids = ConstantList(self._all_token_ids)
|
||||
# trace_headers
|
||||
self.trace_headers = trace_headers
|
||||
self.session_id = session_id
|
||||
|
||||
# True if this request is scheduled as a non-final prefill chunk.
|
||||
self.is_prefill_chunk = False
|
||||
@@ -241,6 +243,7 @@ class Request:
|
||||
trace_headers=request.trace_headers,
|
||||
block_hasher=block_hasher,
|
||||
resumable=request.resumable,
|
||||
session_id=request.session_id,
|
||||
reasoning_ended=request.reasoning_ended,
|
||||
reasoning_parser_kwargs=request.reasoning_parser_kwargs,
|
||||
abort_immediately=request.abort_immediately,
|
||||
|
||||
Reference in New Issue
Block a user