feat: first-class session identity in SGLang (#29436)

This commit is contained in:
ishandhanani
2026-06-28 07:16:38 -07:00
committed by GitHub
parent 828411e6f1
commit aaa31eb0a1
24 changed files with 136 additions and 64 deletions
+2
View File
@@ -68,6 +68,7 @@ message TextGenerateRequest {
optional string routing_key = 10;
optional int32 routed_dp_rank = 11;
map<string, string> trace_headers = 12;
optional string session_id = 13;
}
message TextGenerateResponse {
@@ -90,6 +91,7 @@ message GenerateRequest {
optional string routing_key = 9;
optional int32 routed_dp_rank = 10;
map<string, string> trace_headers = 11;
optional string session_id = 12;
}
message GenerateResponse {
@@ -463,18 +463,21 @@ class Runtime:
self,
prompt: str,
sampling_params: Optional[Dict] = None,
session_id: Optional[str] = None,
):
if self.server_args.skip_tokenizer_init:
json_data = {
"input_ids": prompt,
"sampling_params": sampling_params,
"stream": True,
"session_id": session_id,
}
else:
json_data = {
"text": prompt,
"sampling_params": sampling_params,
"stream": True,
"session_id": session_id,
}
pos = 0
@@ -505,6 +508,7 @@ class Runtime:
logprob_start_len: Optional[Union[List[int], int]] = None,
top_logprobs_num: Optional[Union[List[int], int]] = None,
lora_path: Optional[List[Optional[str]]] = None,
session_id: Optional[str] = None,
):
json_data = {
"text": prompt,
@@ -513,6 +517,7 @@ class Runtime:
"logprob_start_len": logprob_start_len,
"top_logprobs_num": top_logprobs_num,
"lora_path": lora_path,
"session_id": session_id,
}
assert not isinstance(lora_path, list) or len(lora_path) == len(prompt)
response = requests.post(
@@ -33,6 +33,7 @@ class EngineBase(ABC):
data_parallel_rank: Optional[int] = None,
rid: Optional[Union[List[str], str]] = None,
priority: Optional[int] = None,
session_id: Optional[str] = None,
) -> Union[Dict, Iterator[Dict]]:
"""Generate outputs based on given inputs."""
pass
+4
View File
@@ -357,6 +357,7 @@ class Engine(EngineScoreMixin, EngineBase):
rid: Optional[Union[List[str], str]] = None,
session_params: Optional[Dict] = None,
priority: Optional[int] = None,
session_id: Optional[str] = None,
) -> Union[Dict, Iterator[Dict]]:
"""
The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`.
@@ -392,6 +393,7 @@ class Engine(EngineScoreMixin, EngineBase):
disagg_prefill_dp_rank=disagg_prefill_dp_rank,
external_trace_header=external_trace_header,
rid=rid,
session_id=session_id,
session_params=session_params,
priority=priority,
)
@@ -459,6 +461,7 @@ class Engine(EngineScoreMixin, EngineBase):
rid: Optional[Union[List[str], str]] = None,
session_params: Optional[Dict] = None,
priority: Optional[int] = None,
session_id: Optional[str] = None,
) -> Union[Dict, AsyncIterator[Dict]]:
"""
The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`.
@@ -494,6 +497,7 @@ class Engine(EngineScoreMixin, EngineBase):
disagg_prefill_dp_rank=disagg_prefill_dp_rank,
external_trace_header=external_trace_header,
rid=rid,
session_id=session_id,
session_params=session_params,
priority=priority,
)
@@ -115,6 +115,7 @@ class HttpServerEngineAdapter(EngineBase):
lora_path=None,
custom_logit_processor=None,
priority=None,
session_id=None,
):
payload = {
"text": prompt,
@@ -128,6 +129,7 @@ class HttpServerEngineAdapter(EngineBase):
"lora_path": lora_path,
"custom_logit_processor": custom_logit_processor,
"priority": priority,
"session_id": session_id,
}
# Filter out None values
payload = {k: v for k, v in payload.items() if v is not None}
@@ -356,6 +356,7 @@ class CompletionRequest(BaseModel):
ignore_eos: bool = False
skip_special_tokens: bool = True
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
session_id: Optional[str] = None
session_params: Optional[Dict] = None
response_format: Optional[Union[ResponseFormat, StructuralTagResponseFormat]] = None
custom_params: Optional[Dict] = None
@@ -728,6 +729,7 @@ class ChatCompletionRequest(BaseModel):
continue_final_message: bool = False
skip_special_tokens: bool = True
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
session_id: Optional[str] = None
session_params: Optional[Dict] = None
separate_reasoning: bool = True
stream_reasoning: bool = True
@@ -1395,6 +1397,7 @@ class ResponsesRequest(BaseModel):
default_factory=lambda: f"resp_{uuid.uuid4().hex}",
description="The request_id related to this request. If the caller does not set it, a random uuid will be generated.",
)
session_id: Optional[str] = None
priority: int = Field(default=0, description="Request priority")
extra_key: Optional[str] = Field(
default=None,
@@ -606,6 +606,7 @@ class OpenAIServingChat(OpenAIServingBase):
return_routed_experts=request.return_routed_experts,
routed_experts_start_len=request.routed_experts_start_len,
rid=request.rid,
session_id=request.session_id,
extra_key=self._compute_extra_key(request),
require_reasoning=require_reasoning,
priority=request.priority,
@@ -125,6 +125,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
return_routed_experts=request.return_routed_experts,
routed_experts_start_len=request.routed_experts_start_len,
rid=request.rid,
session_id=request.session_id,
extra_key=self._compute_extra_key(request),
priority=request.priority,
routing_key=self.extract_routing_key(raw_request),
@@ -360,6 +360,7 @@ class OpenAIServingResponses(OpenAIServingChat):
sampling_params=sampling_params,
stream=request.stream,
rid=request.request_id,
session_id=request.session_id,
extra_key=self._compute_extra_key(request),
background=request.background,
)
@@ -2362,6 +2363,7 @@ class OpenAIServingResponses(OpenAIServingChat):
sampling_params=sampling_params,
stream=adapted_request.stream,
rid=request_id,
session_id=adapted_request.session_id,
extra_key=adapted_request.extra_key,
return_logprob=adapted_request.return_logprob,
logprob_start_len=adapted_request.logprob_start_len,
+7
View File
@@ -153,6 +153,9 @@ class GenerateReqInput:
# Request ID(s). If omitted, generated during normalization. For batch
# requests, a string is expanded to per-item IDs using it as a prefix.
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
# Stable identity shared by requests in the same session. Unlike
# session_params, this does not alter or reconstruct the prompt.
session_id: Optional[str] = field(default=None, kw_only=True)
# The input prompt. It can be a single prompt or a batch of prompts.
text: Optional[Union[List[str], str]] = None
# The token ids for text.
@@ -338,6 +341,8 @@ class GenerateReqInput:
"""
self._validate_inputs()
self._determine_batch_size()
if self.session_id is not None and self.session_params is not None:
raise ValueError("session_id and session_params cannot both be set.")
self._handle_parallel_sampling()
if self.is_single:
@@ -693,6 +698,7 @@ class GenerateReqInput:
return cache[i]
sub = GenerateReqInput(
rid=self.rid[i],
session_id=self.session_id,
text=self.text[i] if self.text is not None else None,
input_ids=self.input_ids[i] if self.input_ids is not None else None,
input_embeds=(
@@ -800,6 +806,7 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
return_indexer_topk: bool = False
# Session info for continual prompting
session_id: Optional[str] = field(default=None, kw_only=True)
session_params: Optional[SessionParams] = None
# LoRA related
@@ -708,6 +708,7 @@ class Req(ReqDllmMixin):
] = None,
return_pooled_hidden_states: bool = False,
multi_item_delimiter_indices: Optional[List[int]] = None,
session_id: Optional[str] = None,
):
# Input and output info
self.rid = rid
@@ -730,6 +731,7 @@ class Req(ReqDllmMixin):
self.dllm_initialized: bool = False
self.session = session
self.session_id = session_id
self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides
self.multi_item_delimiter_indices = multi_item_delimiter_indices
+9 -43
View File
@@ -123,7 +123,6 @@ from sglang.srt.managers.io_struct import (
LoadLoRAAdapterReqInput,
LoadLoRAAdapterReqOutput,
OpenSessionReqInput,
OpenSessionReqOutput,
PauseGenerationReqInput,
ProfileReq,
ReleaseMemoryOccupationReqInput,
@@ -2011,36 +2010,11 @@ class Scheduler(
session_id = (
recv_req.session_params.id if recv_req.session_params is not None else None
)
# Radix-native session: session_id is just a tag; KV bulk-freed on close.
# Radix-native sessions use only the top-level session_id.
radix_native_session = (
session_id is not None and self.server_args.enable_session_radix_cache
recv_req.session_id is not None
and self.server_args.enable_session_radix_cache
)
if radix_native_session:
sp = recv_req.session_params
if (
sp.rid is not None
or sp.offset is not None
or sp.replace is not None
or sp.drop_previous_output is not None
):
error_msg = (
"Invalid request: radix-native sessions do not support "
"session_params rid/offset/replace/drop_previous_output; "
"send full context each turn."
)
req = Req(
recv_req.rid,
recv_req.input_text,
recv_req.input_ids,
recv_req.sampling_params,
vocab_size=self.model_config.vocab_size,
http_worker_ipc=recv_req.http_worker_ipc,
)
req.tokenizer = self.tokenizer
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if session_id is None or radix_native_session:
# Normal non-session request, or a radix-native session request
@@ -2063,6 +2037,7 @@ class Scheduler(
token_ids_logprob=recv_req.token_ids_logprob,
stream=recv_req.stream,
lora_id=recv_req.lora_id,
session_id=recv_req.session_id,
input_embeds=recv_req.input_embeds,
positional_embed_overrides=recv_req.positional_embed_overrides,
token_type_ids=recv_req.token_type_ids,
@@ -2094,8 +2069,6 @@ class Scheduler(
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
)
req.tokenizer = self.tokenizer
if radix_native_session:
req.session_id = session_id
if self.disaggregation_mode != DisaggregationMode.NULL:
# Invalid request for disaggregated mode
@@ -4082,24 +4055,17 @@ class Scheduler(
return ExpertDistributionReqOutput()
def open_session(self, recv_req: OpenSessionReqInput):
if self.server_args.enable_session_radix_cache:
# Radix-native: open is implicit; explicit open only permits id reuse.
session_id = recv_req.session_id
self.tree_cache.register_session(session_id)
output = OpenSessionReqOutput(
session_id=session_id, success=session_id is not None
)
else:
output = self.session_controller.open(recv_req)
output = self.session_controller.open(recv_req)
if self.ps.pp_rank == 0 and self.ps.tp_rank == 0 and self.ps.attn_cp_rank == 0:
return output
return None
def close_session(self, recv_req: CloseSessionReqInput):
if self.server_args.enable_session_radix_cache:
# "Close" just triggers eviction of the session's tagged KV.
self.tree_cache.release_session(recv_req.session_id)
else:
self.tree_cache.release_radix_session(recv_req.session_id)
if recv_req.session_id in self.session_controller or not (
self.server_args.enable_session_radix_cache
):
self.session_controller.close(recv_req)
def maybe_sleep_on_idle(self):
@@ -1165,6 +1165,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
lora_id=obj.lora_id,
input_embeds=input_embeds,
positional_embed_overrides=obj.positional_embed_overrides,
session_id=obj.session_id,
session_params=session_params,
custom_logit_processor=obj.custom_logit_processor,
require_reasoning=obj.require_reasoning,
@@ -337,6 +337,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
def release_session(self, session_id: str) -> None:
pass
def release_radix_session(self, session_id: str) -> None:
pass
def session_held_tokens(self, active_pool_idxs: Optional[set] = None) -> int:
return 0
@@ -31,6 +31,7 @@ class CacheInitParams:
enable_metrics: bool = False
enable_kv_cache_events: bool = False
enable_session_radix_cache: bool = False
enable_mamba_extra_buffer: bool = False
enable_mamba_extra_buffer_lazy: bool = False
@@ -219,6 +219,7 @@ def build_kv_cache(
eviction_policy=server_args.radix_eviction_policy,
enable_metrics=enable_metrics,
enable_kv_cache_events=enable_kv_cache_events,
enable_session_radix_cache=server_args.enable_session_radix_cache,
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(),
pp_rank=ps.pp_rank,
@@ -290,6 +290,7 @@ class RadixCache(SessionRadixCacheMixin, KVCacheEventMixin, BasePrefixCache):
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
self.page_size = params.page_size
self.enable_kv_cache_events = params.enable_kv_cache_events
self.enable_session_radix_cache = params.enable_session_radix_cache
self.is_eagle = params.is_eagle
self.disable_finished_insert = params.disable_finished_insert
self.eviction_policy = params.eviction_policy.lower()
@@ -34,13 +34,6 @@ class SessionRadixCacheMixin:
if not hasattr(self, "_session_leaves"):
self._reset_session_radix_state()
def register_session(self, session_id: str) -> None:
self._ensure_session_radix_state()
if session_id is None:
return
self._closed_session_ids.pop(session_id, None)
self._session_leaves.setdefault(session_id, set())
def _remember_closed_session(self, session_id: str) -> None:
self._closed_session_ids[session_id] = None
self._closed_session_ids.move_to_end(session_id)
@@ -62,6 +55,8 @@ class SessionRadixCacheMixin:
def _tag_session_leaf(self, req: Req, radix_key, node=None) -> None:
"""Add this request's session id to its leaf's holder set; no-op for non-session reqs."""
if not self.enable_session_radix_cache:
return
self._ensure_session_radix_state()
sid = getattr(req, "session_id", None)
if sid is None or sid in self._closed_session_ids:
@@ -86,7 +81,7 @@ class SessionRadixCacheMixin:
len(self._session_leaves[sid]),
)
def release_session(self, session_id: str) -> int:
def release_radix_session(self, session_id: str) -> int:
"""Close: drop this session from each of its tagged leaves, freeing a node
only once no other session still holds it (last holder). Shared
prefixes/leaves kept."""
@@ -438,6 +438,9 @@ class StreamingSession(BasePrefixCache):
self._free_slot_mamba(slot)
def release_radix_session(self, session_id: str) -> None:
self.inner.release_radix_session(session_id)
def session_held_tokens(self, active_pool_idxs: Optional[set] = None) -> int:
"""Total KV tokens held by session slots, not tracked by the tree.
@@ -128,6 +128,9 @@ pub(crate) fn build_text_generate_dict(
if let Some(rank) = req.routed_dp_rank {
d.insert("routed_dp_rank".into(), serde_json::json!(rank));
}
if let Some(ref session_id) = req.session_id {
d.insert("session_id".into(), serde_json::json!(session_id));
}
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
d.insert("external_trace_header".into(), trace);
}
@@ -172,6 +175,9 @@ pub(crate) fn build_generate_dict(
if let Some(rank) = req.routed_dp_rank {
d.insert("routed_dp_rank".into(), serde_json::json!(rank));
}
if let Some(ref session_id) = req.session_id {
d.insert("session_id".into(), serde_json::json!(session_id));
}
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
d.insert("external_trace_header".into(), trace);
}
@@ -237,3 +243,30 @@ pub(crate) fn build_classify_dict(
d.insert("received_time".into(), serde_json::json!(now_timestamp()));
d
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn generate_dicts_include_session_id() {
let session_id = Some("session-1".to_string());
let text_req = proto::TextGenerateRequest {
session_id: session_id.clone(),
..Default::default()
};
let token_req = proto::GenerateRequest {
session_id,
..Default::default()
};
assert_eq!(
build_text_generate_dict("request-1", &text_req).get("session_id"),
Some(&serde_json::json!("session-1"))
);
assert_eq!(
build_generate_dict("request-2", &token_req).get("session_id"),
Some(&serde_json::json!("session-1"))
);
}
}
+24 -12
View File
@@ -46,6 +46,7 @@ class TestSessionRadixCache(unittest.TestCase):
page_size=1,
eviction_policy="lru",
enable_kv_cache_events=False,
enable_session_radix_cache=True,
)
)
@@ -94,13 +95,19 @@ class TestSessionRadixCache(unittest.TestCase):
self.assertEqual(getattr(leaf, "session_ids", None), {"S"})
self.assertIn(leaf, self.cache._session_leaves["S"])
def test_disabled_cache_does_not_tag_session_kv(self):
self.cache.enable_session_radix_cache = False
self._insert([1, 2, 3, 4])
self._tag([1, 2, 3, 4], "S")
self.assertIsNone(getattr(self._leaf([1, 2, 3, 4]), "session_ids", None))
def test_shared_prefix_frees_only_unique_tail(self):
# A/B share prefix [1,2]; close(A) frees only A's tail, B + shared stay.
self._insert([1, 2, 3, 4])
self._tag([1, 2, 3, 4], "A")
self._insert([1, 2, 5, 6])
self._tag([1, 2, 5, 6], "B")
self.assertGreater(self.cache.release_session("A"), 0)
self.assertGreater(self.cache.release_radix_session("A"), 0)
self.assertEqual(self._cached([1, 2, 3, 4]), 2) # only shared [1,2] left
self.assertEqual(self._cached([1, 2, 5, 6]), 4) # B intact
@@ -113,11 +120,18 @@ class TestSessionRadixCache(unittest.TestCase):
self.assertEqual(
getattr(self._leaf([1, 2, 3, 4]), "session_ids", None), {"A", "B"}
)
self.assertEqual(self.cache.release_session("A"), 0) # B still holds
self.assertEqual(self.cache.release_radix_session("A"), 0) # B still holds
self.assertEqual(self._cached([1, 2, 3, 4]), 4)
self.assertEqual(self.cache.release_session("B"), 1) # last holder frees
self.assertEqual(self.cache.release_radix_session("B"), 1) # last holder frees
self.assertEqual(self._cached([1, 2, 3, 4]), 0)
def test_legacy_release_does_not_release_radix_session(self):
self._insert([1, 2, 3, 4])
self._tag([1, 2, 3, 4], "S")
self.cache.release_session("S")
self.assertEqual(self._cached([1, 2, 3, 4]), 4)
self.assertEqual(self.cache.release_radix_session("S"), 1)
def test_tag_is_lru_neutral_not_pinned(self):
# The tag must add no lock/pin: a tagged, never-closed node is evictable.
self._insert([1, 2, 3, 4])
@@ -129,29 +143,27 @@ class TestSessionRadixCache(unittest.TestCase):
self.cache.evict(EvictParams(num_tokens=4)) # LRU reclaims it while open
self.assertEqual(self._cached([1, 2, 3, 4]), 0)
self.assertNotIn("S", self.cache._session_leaves)
self.assertEqual(self.cache.release_session("S"), 0) # late close is a no-op
self.assertEqual(
self.cache.release_radix_session("S"), 0
) # late close is a no-op
def test_close_tombstone_blocks_late_finish_until_reopen(self):
def test_close_tombstone_blocks_late_finish(self):
self._insert([1, 2, 3, 4])
self._tag([1, 2, 3, 4], "S")
self.assertEqual(self.cache.release_session("S"), 1)
self.assertEqual(self.cache.release_radix_session("S"), 1)
self._insert([5, 6, 7, 8])
self._tag([5, 6, 7, 8], "S") # simulates a finish racing after close
self.assertIsNone(getattr(self._leaf([5, 6, 7, 8]), "session_ids", None))
self.cache.register_session("S")
self._tag([5, 6, 7, 8], "S")
self.assertEqual(getattr(self._leaf([5, 6, 7, 8]), "session_ids", None), {"S"})
def test_tombstoned_shared_holder_cannot_retag_after_last_holder_close(self):
self._insert([1, 2, 3, 4])
self._tag([1, 2, 3, 4], "A")
self._tag([1, 2, 3, 4], "B")
self.assertEqual(self.cache.release_session("B"), 0)
self.assertEqual(self.cache.release_radix_session("B"), 0)
self.assertEqual(getattr(self._leaf([1, 2, 3, 4]), "session_ids", None), {"A"})
self.assertEqual(self.cache.release_session("A"), 1)
self.assertEqual(self.cache.release_radix_session("A"), 1)
self._insert([5, 6, 7, 8])
self._tag([5, 6, 7, 8], "B")
@@ -102,6 +102,7 @@ class ServingChatTestCase(unittest.TestCase):
self.basic_req = ChatCompletionRequest(
model="x",
messages=[{"role": "user", "content": "Hi?"}],
session_id="session-1",
temperature=0.7,
max_tokens=100,
stream=False,
@@ -145,6 +146,7 @@ class ServingChatTestCase(unittest.TestCase):
adapted, processed = self.chat._convert_to_internal_request(self.basic_req)
self.assertIsInstance(adapted, GenerateReqInput)
self.assertFalse(adapted.stream)
self.assertEqual(adapted.session_id, "session-1")
self.assertEqual(processed, self.basic_req)
def test_convert_to_internal_request_rejects_stream_return_prompt_token_ids(self):
@@ -61,9 +61,15 @@ class ServingCompletionTestCase(unittest.TestCase):
# ---------- prompt-handling ----------
def test_single_string_prompt(self):
req = CompletionRequest(model="x", prompt="Hello world", max_tokens=100)
req = CompletionRequest(
model="x",
prompt="Hello world",
max_tokens=100,
session_id="session-1",
)
internal, _ = self.sc._convert_to_internal_request(req)
self.assertEqual(internal.text, "Hello world")
self.assertEqual(internal.session_id, "session-1")
def test_single_token_ids_prompt(self):
req = CompletionRequest(model="x", prompt=[1, 2, 3, 4], max_tokens=100)
@@ -551,6 +551,24 @@ class TestGenerateReqInputNormalization(CustomTestCase):
req.normalize_batch_and_arguments()
self.assertEqual(req.session_params, [{"id": "session1"}, {"id": "session2"}])
def test_session_id_handling(self):
req = GenerateReqInput(
text=["Hello", "World"],
session_id="session1",
sampling_params={"n": 2},
)
req.normalize_batch_and_arguments()
self.assertEqual(req.session_id, "session1")
self.assertIsNone(req.session_params)
self.assertEqual(req[2].session_id, "session1")
with self.assertRaisesRegex(ValueError, "cannot both be set"):
GenerateReqInput(
text="Hello",
session_id="explicit",
session_params={"id": "legacy"},
).normalize_batch_and_arguments()
def test_getitem_method(self):
"""Test the __getitem__ method."""
req = GenerateReqInput(