diff --git a/rust/proto/inference.proto b/rust/proto/inference.proto index b1c08ae76e6..575ba70a19e 100644 --- a/rust/proto/inference.proto +++ b/rust/proto/inference.proto @@ -45,6 +45,8 @@ message GenerateRequest { uint32 truncate_prompt_tokens = 11; int32 priority = 12; + + optional string session_id = 13; } message RandomSampling { diff --git a/rust/src/chat/src/lib.rs b/rust/src/chat/src/lib.rs index 0d538a24ba1..eb40272604e 100644 --- a/rust/src/chat/src/lib.rs +++ b/rust/src/chat/src/lib.rs @@ -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), diff --git a/rust/src/chat/src/request.rs b/rust/src/chat/src/request.rs index 56a764f8987..f4cc252d81e 100644 --- a/rust/src/chat/src/request.rs +++ b/rust/src/chat/src/request.rs @@ -491,6 +491,9 @@ pub struct ChatRequest { /// Override data parallel rank. #[serde(default)] pub data_parallel_rank: Option, + /// Stable session identity shared by related requests. + #[serde(default)] + pub session_id: Option, /// LoRA adapter selected for this request. #[serde(default)] pub lora_request: Option, @@ -514,6 +517,7 @@ impl ChatRequest { cache_salt: None, add_special_tokens: false, data_parallel_rank: None, + session_id: None, lora_request: None, } } diff --git a/rust/src/engine-core-client/src/protocol/request.rs b/rust/src/engine-core-client/src/protocol/request.rs index 99e1b588e17..8118b770e3c 100644 --- a/rust/src/engine-core-client/src/protocol/request.rs +++ b/rust/src/engine-core-client/src/protocol/request.rs @@ -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, } 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] diff --git a/rust/src/engine-core-client/src/tests/client.rs b/rust/src/engine-core-client/src/tests/client.rs index 11d08db646b..3268190785b 100644 --- a/rust/src/engine-core-client/src/tests/client.rs +++ b/rust/src/engine-core-client/src/tests/client.rs @@ -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() } } diff --git a/rust/src/engine-core-client/src/tests/python_compat.py b/rust/src/engine-core-client/src/tests/python_compat.py index c68e2b6b471..0d95dbba195 100755 --- a/rust/src/engine-core-client/src/tests/python_compat.py +++ b/rust/src/engine-core-client/src/tests/python_compat.py @@ -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. diff --git a/rust/src/llm/examples/external_engine_smoke.rs b/rust/src/llm/examples/external_engine_smoke.rs index ae4f428822c..9d318285ae8 100644 --- a/rust/src/llm/examples/external_engine_smoke.rs +++ b/rust/src/llm/examples/external_engine_smoke.rs @@ -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, } diff --git a/rust/src/llm/src/request.rs b/rust/src/llm/src/request.rs index c4d78396bac..c647f54d8df 100644 --- a/rust/src/llm/src/request.rs +++ b/rust/src/llm/src/request.rs @@ -46,6 +46,8 @@ pub struct GenerateRequest { pub priority: i32, /// Optional data-parallel rank override for routing this request. pub data_parallel_rank: Option, + /// Stable session identity shared by related requests. + pub session_id: Option, /// Optional reasoning-parser kwargs forwarded to engine-side structured /// output logic. pub reasoning_parser_kwargs: Option, @@ -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([( diff --git a/rust/src/llm/tests/generate.rs b/rust/src/llm/tests/generate.rs index d6bc9d1cfaa..5e5d599cd97 100644 --- a/rust/src/llm/tests/generate.rs +++ b/rust/src/llm/tests/generate.rs @@ -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, } diff --git a/rust/src/server/src/grpc/convert.rs b/rust/src/server/src/grpc/convert.rs index 2f13422fecc..178da515ee8 100644 --- a/rust/src/server/src/grpc/convert.rs +++ b/rust/src/server/src/grpc/convert.rs @@ -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, diff --git a/rust/src/server/src/routes/inference/generate/convert.rs b/rust/src/server/src/routes/inference/generate/convert.rs index 02a52b04146..10361e679ad 100644 --- a/rust/src/server/src/routes/inference/generate/convert.rs +++ b/rust/src/server/src/routes/inference/generate/convert.rs @@ -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, diff --git a/rust/src/server/src/routes/openai/chat_completions/convert.rs b/rust/src/server/src/routes/openai/chat_completions/convert.rs index 2c4050f8f1e..030f50f9721 100644 --- a/rust/src/server/src/routes/openai/chat_completions/convert.rs +++ b/rust/src/server/src/routes/openai/chat_completions/convert.rs @@ -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( diff --git a/rust/src/server/src/routes/openai/chat_completions/types.rs b/rust/src/server/src/routes/openai/chat_completions/types.rs index 1ab2f4c604c..c0735553c5a 100644 --- a/rust/src/server/src/routes/openai/chat_completions/types.rs +++ b/rust/src/server/src/routes/openai/chat_completions/types.rs @@ -222,6 +222,9 @@ pub struct ChatCompletionRequest { /// External request ID used for response correlation. pub request_id: Option, + /// Stable session identity shared by related requests. + pub session_id: Option, + /// Tokens represented as strings of the form 'token_id:{token_id}' in /// logprobs pub return_tokens_as_token_ids: Option, @@ -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, diff --git a/rust/src/server/src/routes/openai/completions/convert.rs b/rust/src/server/src/routes/openai/completions/convert.rs index 2030a4ad408..1143da77f27 100644 --- a/rust/src/server/src/routes/openai/completions/convert.rs +++ b/rust/src/server/src/routes/openai/completions/convert.rs @@ -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!({ diff --git a/rust/src/server/src/routes/openai/completions/types.rs b/rust/src/server/src/routes/openai/completions/types.rs index 09d40215012..504c3e3b2c2 100644 --- a/rust/src/server/src/routes/openai/completions/types.rs +++ b/rust/src/server/src/routes/openai/completions/types.rs @@ -164,6 +164,9 @@ pub struct CompletionRequest { /// External request ID used for response correlation. pub request_id: Option, + /// Stable session identity shared by related requests. + pub session_id: Option, + /// Tokens represented as strings of the form 'token_id:{token_id}' in /// logprobs pub return_tokens_as_token_ids: Option, diff --git a/rust/src/server/src/routes/tokenize/types.rs b/rust/src/server/src/routes/tokenize/types.rs index 6684aa55d26..39dbeb49a2d 100644 --- a/rust/src/server/src/routes/tokenize/types.rs +++ b/rust/src/server/src/routes/tokenize/types.rs @@ -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, }) } diff --git a/rust/src/server/src/utils.rs b/rust/src/server/src/utils.rs index 8879e0ef89d..2c9e0e51129 100644 --- a/rust/src/server/src/utils.rs +++ b/rust/src/server/src/utils.rs @@ -15,6 +15,7 @@ use crate::error::ApiError; pub struct ResolvedRequestContext { pub request_id: String, pub data_parallel_rank: Option, + pub session_id: Option, } /// 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>, +) -> Option { + 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` 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, } } diff --git a/rust/src/text/src/lower.rs b/rust/src/text/src/lower.rs index aa54595acda..9c0a5905d54 100644 --- a/rust/src/text/src/lower.rs +++ b/rust/src/text/src/lower.rs @@ -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, diff --git a/rust/src/text/src/request.rs b/rust/src/text/src/request.rs index 8a8dc0e7d94..dfe7dc742cc 100644 --- a/rust/src/text/src/request.rs +++ b/rust/src/text/src/request.rs @@ -183,6 +183,9 @@ pub struct TextRequest { /// Override data parallel rank. #[serde(default)] pub data_parallel_rank: Option, + /// Stable session identity shared by related requests. + #[serde(default)] + pub session_id: Option, /// 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, diff --git a/tests/entrypoints/openai/test_session_id.py b/tests/entrypoints/openai/test_session_id.py new file mode 100644 index 00000000000..d8f68c43ef8 --- /dev/null +++ b/tests/entrypoints/openai/test_session_id.py @@ -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 diff --git a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py index ce3100f196c..8936b19093b 100644 --- a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py +++ b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py @@ -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].""" diff --git a/tests/v1/engine/test_parallel_sampling.py b/tests/v1/engine/test_parallel_sampling.py index 395867c0600..72e3e4e351d 100644 --- a/tests/v1/engine/test_parallel_sampling.py +++ b/tests/v1/engine/test_parallel_sampling.py @@ -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", diff --git a/tests/v1/test_request.py b/tests/v1/test_request.py index 66e46faa7ad..be417b9b2ff 100644 --- a/tests/v1/test_request.py +++ b/tests/v1/test_request.py @@ -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" diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index 5a9b9f96d2c..0f61b73839f 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -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]: diff --git a/vllm/entrypoints/generate/base/serving.py b/vllm/entrypoints/generate/base/serving.py index 78f84ca872b..2497f2f05cf 100644 --- a/vllm/entrypoints/generate/base/serving.py +++ b/vllm/entrypoints/generate/base/serving.py @@ -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], diff --git a/vllm/entrypoints/generate/beam_search/online.py b/vllm/entrypoints/generate/beam_search/online.py index 1cd821f9db8..02ff61d6293 100644 --- a/vllm/entrypoints/generate/beam_search/online.py +++ b/vllm/entrypoints/generate/beam_search/online.py @@ -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, ) ) ) diff --git a/vllm/entrypoints/openai/chat_completion/batch_serving.py b/vllm/entrypoints/openai/chat_completion/batch_serving.py index 17fff37cc13..275ad13f06e 100644 --- a/vllm/entrypoints/openai/chat_completion/batch_serving.py +++ b/vllm/entrypoints/openai/chat_completion/batch_serving.py @@ -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, ) ) diff --git a/vllm/entrypoints/openai/chat_completion/protocol.py b/vllm/entrypoints/openai/chat_completion/protocol.py index fc3b8103bc1..d25b3886ad5 100644 --- a/vllm/entrypoints/openai/chat_completion/protocol.py +++ b/vllm/entrypoints/openai/chat_completion/protocol.py @@ -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, diff --git a/vllm/entrypoints/openai/chat_completion/serving.py b/vllm/entrypoints/openai/chat_completion/serving.py index 7b3ce551b20..3ab4c1ceaa3 100644 --- a/vllm/entrypoints/openai/chat_completion/serving.py +++ b/vllm/entrypoints/openai/chat_completion/serving.py @@ -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, diff --git a/vllm/entrypoints/openai/completion/protocol.py b/vllm/entrypoints/openai/completion/protocol.py index 31c35b5bcce..516a2c371ff 100644 --- a/vllm/entrypoints/openai/completion/protocol.py +++ b/vllm/entrypoints/openai/completion/protocol.py @@ -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, diff --git a/vllm/entrypoints/openai/completion/serving.py b/vllm/entrypoints/openai/completion/serving.py index 8791042c692..7c47e777bff 100644 --- a/vllm/entrypoints/openai/completion/serving.py +++ b/vllm/entrypoints/openai/completion/serving.py @@ -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) diff --git a/vllm/entrypoints/openai/responses/protocol.py b/vllm/entrypoints/openai/responses/protocol.py index c739ffe25c3..ec8d25d10a6 100644 --- a/vllm/entrypoints/openai/responses/protocol.py +++ b/vllm/entrypoints/openai/responses/protocol.py @@ -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=( diff --git a/vllm/entrypoints/openai/responses/serving.py b/vllm/entrypoints/openai/responses/serving.py index 344ce0a209c..44b78029579 100644 --- a/vllm/entrypoints/openai/responses/serving.py +++ b/vllm/entrypoints/openai/responses/serving.py @@ -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, ) diff --git a/vllm/entrypoints/scale_out/token_in_token_out/serving.py b/vllm/entrypoints/scale_out/token_in_token_out/serving.py index 34e9eaeb12d..7b5e2bbb3d9 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/serving.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/serving.py @@ -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 diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 51428208022..0197438e238 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -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).""" diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index 4ce9d64b440..39500941e28 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -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, diff --git a/vllm/v1/engine/input_processor.py b/vllm/v1/engine/input_processor.py index 1fe460b54af..359d1952a10 100644 --- a/vllm/v1/engine/input_processor.py +++ b/vllm/v1/engine/input_processor.py @@ -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( diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py index 17e40630859..d015fb16d35 100644 --- a/vllm/v1/engine/llm_engine.py +++ b/vllm/v1/engine/llm_engine.py @@ -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) diff --git a/vllm/v1/request.py b/vllm/v1/request.py index 4a2cc8d8dbc..0b969c991d9 100644 --- a/vllm/v1/request.py +++ b/vllm/v1/request.py @@ -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,