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
@@ -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.