feat(frontend): session id plumbing into requests (#48048)

Signed-off-by: Karen Chung <[email protected]>
This commit is contained in:
Karen Chung
2026-08-03 22:18:58 +00:00
committed by GitHub
parent e578de311c
commit f57123aa2d
39 changed files with 438 additions and 4 deletions
+2
View File
@@ -45,6 +45,8 @@ message GenerateRequest {
uint32 truncate_prompt_tokens = 11;
int32 priority = 12;
optional string session_id = 13;
}
message RandomSampling {
+1
View File
@@ -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),
+4
View File
@@ -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,
}
+6
View File
@@ -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([(
+1
View File
@@ -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,
}
+2
View File
@@ -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,
})
}
+27 -2
View File
@@ -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,
}
}
+1
View File
@@ -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,
+4
View File
@@ -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]."""
+16
View File
@@ -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",
+23 -1
View File
@@ -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"
+1
View File
@@ -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]:
+24
View File
@@ -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
+2
View File
@@ -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)."""
+8
View File
@@ -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,
+2
View File
@@ -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(
+2
View File
@@ -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)
+3
View File
@@ -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,