[Session] Extract SessionController and clean up session logic in Scheduler (#19547)

This commit is contained in:
Liangsheng Yin
2026-02-28 19:47:44 -08:00
committed by GitHub
parent a45613f2a6
commit 5acb45cf32
2 changed files with 142 additions and 83 deletions
+38 -80
View File
@@ -22,7 +22,7 @@ import time
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import Any, Deque, Dict, List, Optional, Tuple, Union from typing import Any, Deque, List, Optional, Tuple, Union
import psutil import psutil
import setproctitle import setproctitle
@@ -112,7 +112,6 @@ from sglang.srt.managers.io_struct import (
LoadLoRAAdapterReqInput, LoadLoRAAdapterReqInput,
LoadLoRAAdapterReqOutput, LoadLoRAAdapterReqOutput,
OpenSessionReqInput, OpenSessionReqInput,
OpenSessionReqOutput,
PauseGenerationReqInput, PauseGenerationReqInput,
ProfileReq, ProfileReq,
ReleaseMemoryOccupationReqInput, ReleaseMemoryOccupationReqInput,
@@ -167,7 +166,7 @@ from sglang.srt.managers.scheduler_runtime_checker_mixin import (
from sglang.srt.managers.scheduler_update_weights_mixin import ( from sglang.srt.managers.scheduler_update_weights_mixin import (
SchedulerUpdateWeightsMixin, SchedulerUpdateWeightsMixin,
) )
from sglang.srt.managers.session_controller import Session from sglang.srt.managers.session_controller import SessionController
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import release_kv_cache from sglang.srt.mem_cache.common import release_kv_cache
@@ -755,8 +754,7 @@ class Scheduler(
self.return_health_check_ct = 0 self.return_health_check_ct = 0
self.num_retracted_reqs: int = 0 self.num_retracted_reqs: int = 0
self.num_paused_reqs: int = 0 self.num_paused_reqs: int = 0
self.sessions: Dict[str, Session] = {} self.session_controller = SessionController(self.tree_cache)
self._last_reap_sessions: float = 0.0
self.forward_sleep_time = None self.forward_sleep_time = None
self._engine_paused = False self._engine_paused = False
@@ -1131,7 +1129,6 @@ class Scheduler(
self.process_batch_result(batch, result) self.process_batch_result(batch, result)
else: else:
# When the server is idle, do self-check and re-init some states. # When the server is idle, do self-check and re-init some states.
# Skip if there are any streaming sessions (latency sensitive).
self.self_check_during_idle() self.self_check_during_idle()
# Update last_batch # Update last_batch
@@ -1363,9 +1360,7 @@ class Scheduler(
def process_input_requests(self, recv_reqs: List): def process_input_requests(self, recv_reqs: List):
now = time.monotonic() now = time.monotonic()
if now - self._last_reap_sessions > 1.0: # reap sessions every second self.session_controller.maybe_reap(now)
self._last_reap_sessions = now
self.reap_timed_out_sessions()
for recv_req in recv_reqs: for recv_req in recv_reqs:
# If it is a health check generation request and there are running requests, ignore it. # If it is a health check generation request and there are running requests, ignore it.
if is_health_check_generate_req(recv_req) and ( if is_health_check_generate_req(recv_req) and (
@@ -1485,12 +1480,13 @@ class Scheduler(
self, self,
recv_req: TokenizedGenerateReqInput, recv_req: TokenizedGenerateReqInput,
): ):
# Create a new request # Route: normal request / session request / session-not-found
if ( session_id = (
recv_req.session_params is None recv_req.session_params.id if recv_req.session_params is not None else None
or recv_req.session_params.id is None )
or recv_req.session_params.id not in self.sessions
): if session_id is None:
# Normal non-session request
if recv_req.input_embeds is not None: if recv_req.input_embeds is not None:
# Generate fake input_ids based on the length of input_embeds # Generate fake input_ids based on the length of input_embeds
seq_length = len(recv_req.input_embeds) seq_length = len(recv_req.input_embeds)
@@ -1553,21 +1549,14 @@ class Scheduler(
self.stream_output([req], req.return_logprob) self.stream_output([req], req.return_logprob)
return return
if ( elif session_id in self.session_controller:
recv_req.session_params is not None # Session exists: create request from session
and recv_req.session_params.id is not None session = self.session_controller.get(session_id)
):
req.set_finish_with_abort(
f"Invalid request: session id {recv_req.session_params.id} does not exist"
)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
else:
# Create a new request from a previous session
session = self.sessions[recv_req.session_params.id]
req = session.create_req( req = session.create_req(
recv_req, self.tokenizer, self.model_config.vocab_size recv_req,
self.tokenizer,
self.model_config.vocab_size,
eos_token_ids=self.model_config.hf_eos_token_id,
) )
# TODO: set trace context # TODO: set trace context
if self.enable_metrics: if self.enable_metrics:
@@ -1577,22 +1566,28 @@ class Scheduler(
self._add_request_to_queue(req) self._add_request_to_queue(req)
return return
else:
# Session ID provided but session not found
req = Req(
recv_req.rid,
recv_req.input_text,
recv_req.input_ids,
recv_req.sampling_params,
vocab_size=self.model_config.vocab_size,
)
req.tokenizer = self.tokenizer
req.set_finish_with_abort(
f"Invalid request: session id {session_id} does not exist"
)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
# Handle multimodal inputs # Handle multimodal inputs
if recv_req.mm_inputs is not None: if recv_req.mm_inputs is not None:
image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs) image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs)
# For session requests, adjust mm_inputs offsets by the prefix length. SessionController.adjust_mm_offsets(recv_req, req, image_inputs)
# Session.create_req prepends previous context to origin_input_ids,
# so offsets from the new prompt need to be shifted.
if len(recv_req.input_ids) < len(req.origin_input_ids):
assert recv_req.session_params.id in self.sessions
prefix_len = len(req.origin_input_ids) - len(recv_req.input_ids)
for mm_item in image_inputs.mm_items:
if mm_item.offsets:
mm_item.offsets = [
(start + prefix_len, end + prefix_len)
for start, end in mm_item.offsets
]
# The following steps are already fast, execute locally on each rank. # The following steps are already fast, execute locally on each rank.
# Expand a single image token into multiple dummy tokens for receiving image embeddings # Expand a single image token into multiple dummy tokens for receiving image embeddings
@@ -2938,47 +2933,10 @@ class Scheduler(
return ExpertDistributionReqOutput() return ExpertDistributionReqOutput()
def open_session(self, recv_req: OpenSessionReqInput): def open_session(self, recv_req: OpenSessionReqInput):
session_id = recv_req.session_id return self.session_controller.open(recv_req)
if session_id in self.sessions:
logger.warning(f"session id {session_id} already exist, cannot open.")
return OpenSessionReqOutput(session_id, False)
elif session_id is None:
logger.warning("session id is None, cannot open.")
return OpenSessionReqOutput(session_id, False)
else:
self.sessions[session_id] = Session(
recv_req.capacity_of_str_len,
session_id,
streaming=bool(recv_req.streaming),
timeout=recv_req.timeout,
)
return OpenSessionReqOutput(session_id, True)
def close_session(self, recv_req: CloseSessionReqInput): def close_session(self, recv_req: CloseSessionReqInput):
session_id = recv_req.session_id self.session_controller.close(recv_req)
if session_id not in self.sessions:
logger.warning(f"session id {session_id} does not exist, cannot delete.")
else:
self._close_session(session_id)
def _close_session(self, session_id: str):
session = self.sessions[session_id]
if session.streaming and session.req_nodes:
assert len(session.req_nodes) == 1
req = next(iter(session.req_nodes.values())).req
if not req.finished():
req.session = None
if isinstance(self.tree_cache, SessionAwareCache):
self.tree_cache.release_session(session_id)
del self.sessions[session_id]
def reap_timed_out_sessions(self):
timed_out = [
sid for sid, session in self.sessions.items() if session.is_timed_out()
]
for sid in timed_out:
logger.info(f"Session {sid} timed out, closing.")
self._close_session(sid)
def maybe_sleep_on_idle(self): def maybe_sleep_on_idle(self):
if self.idle_sleeper is not None: if self.idle_sleeper is not None:
@@ -10,13 +10,26 @@
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
from __future__ import annotations
import logging import logging
import time import time
import uuid import uuid
from typing import Dict, Optional from typing import TYPE_CHECKING, Dict, Optional
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput from sglang.srt.managers.io_struct import (
CloseSessionReqInput,
OpenSessionReqInput,
OpenSessionReqOutput,
TokenizedGenerateReqInput,
)
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
if TYPE_CHECKING:
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
logger = logging.getLogger(__name__)
class SessionReqNode: class SessionReqNode:
@@ -85,7 +98,13 @@ class Session:
return False return False
return time.monotonic() - self.last_active_time > self.timeout return time.monotonic() - self.last_active_time > self.timeout
def create_req(self, req: TokenizedGenerateReqInput, tokenizer, vocab_size: int): def create_req(
self,
req: TokenizedGenerateReqInput,
tokenizer,
vocab_size: int,
eos_token_ids=None,
):
assert req.session_params is not None assert req.session_params is not None
self.last_active_time = time.monotonic() self.last_active_time = time.monotonic()
session_params = req.session_params session_params = req.session_params
@@ -190,6 +209,14 @@ class Session:
top_logprobs_num=req.top_logprobs_num, top_logprobs_num=req.top_logprobs_num,
token_ids_logprob=req.token_ids_logprob, token_ids_logprob=req.token_ids_logprob,
vocab_size=vocab_size, vocab_size=vocab_size,
eos_token_ids=eos_token_ids,
require_reasoning=req.require_reasoning,
return_hidden_states=req.return_hidden_states,
return_routed_experts=req.return_routed_experts,
priority=req.priority,
routing_key=req.routing_key,
http_worker_ipc=req.http_worker_ipc,
time_stats=req.time_stats,
) )
if last_req is not None: if last_req is not None:
new_req.multimodal_inputs = last_req.multimodal_inputs new_req.multimodal_inputs = last_req.multimodal_inputs
@@ -206,3 +233,77 @@ class Session:
self.req_nodes[req.rid] = new_req_node self.req_nodes[req.rid] = new_req_node
return new_req return new_req
class SessionController:
def __init__(self, tree_cache: BasePrefixCache):
self.sessions: Dict[str, Session] = {}
self._last_reap_time: float = 0.0
self.tree_cache = tree_cache
def __contains__(self, session_id: str) -> bool:
return session_id in self.sessions
def get(self, session_id: str) -> Optional[Session]:
return self.sessions.get(session_id)
def open(self, recv_req: OpenSessionReqInput) -> OpenSessionReqOutput:
session_id = recv_req.session_id
if session_id in self.sessions:
logger.warning(f"session id {session_id} already exist, cannot open.")
return OpenSessionReqOutput(session_id, False)
elif session_id is None:
logger.warning("session id is None, cannot open.")
return OpenSessionReqOutput(session_id, False)
else:
self.sessions[session_id] = Session(
recv_req.capacity_of_str_len,
session_id,
streaming=bool(recv_req.streaming),
timeout=recv_req.timeout,
)
return OpenSessionReqOutput(session_id, True)
def close(self, recv_req: CloseSessionReqInput):
session_id = recv_req.session_id
if session_id not in self.sessions:
logger.warning(f"session id {session_id} does not exist, cannot delete.")
else:
self._close(session_id)
def _close(self, session_id: str):
session = self.sessions[session_id]
if session.streaming and session.req_nodes:
assert len(session.req_nodes) == 1
req = next(iter(session.req_nodes.values())).req
if not req.finished():
req.session = None
if isinstance(self.tree_cache, SessionAwareCache):
self.tree_cache.release_session(session_id)
del self.sessions[session_id]
def maybe_reap(self, now: float, interval: float = 1.0):
# reap sessions every second
if now - self._last_reap_time > interval:
self._last_reap_time = now
timed_out = [
sid for sid, session in self.sessions.items() if session.is_timed_out()
]
for sid in timed_out:
logger.info(f"Session {sid} timed out, closing.")
self._close(sid)
@staticmethod
def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs):
# For session requests, adjust mm_inputs offsets by the prefix length.
# Session.create_req prepends previous context to origin_input_ids,
# so offsets from the new prompt need to be shifted.
if len(recv_req.input_ids) >= len(req.origin_input_ids):
return
prefix_len = len(req.origin_input_ids) - len(recv_req.input_ids)
for mm_item in image_inputs.mm_items:
if mm_item.offsets:
mm_item.offsets = [
(start + prefix_len, end + prefix_len)
for start, end in mm_item.offsets
]