[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

View File

@@ -10,13 +10,26 @@
# limitations under the License.
# ==============================================================================
from __future__ import annotations
import logging
import time
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.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:
@@ -85,7 +98,13 @@ class Session:
return False
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
self.last_active_time = time.monotonic()
session_params = req.session_params
@@ -190,6 +209,14 @@ class Session:
top_logprobs_num=req.top_logprobs_num,
token_ids_logprob=req.token_ids_logprob,
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:
new_req.multimodal_inputs = last_req.multimodal_inputs
@@ -206,3 +233,77 @@ class Session:
self.req_nodes[req.rid] = new_req_node
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
]