fix: streaming session race condition + some metrics (#21875)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
ishandhanani
2026-04-12 18:05:23 -07:00
committed by GitHub
co-authored by Claude Opus 4.6 hnyls2002 Liangsheng Yin
parent 37fc47c645
commit c1ab68b45e
11 changed files with 966 additions and 26 deletions
+17 -7
View File
@@ -1848,8 +1848,11 @@ class Scheduler(
self.stream_output([req], req.return_logprob) self.stream_output([req], req.return_logprob)
return return
elif session_id in self.session_controller: elif (
# Session exists: create request from session session_id in self.session_controller
and not self.session_controller.get(session_id).close_on_finish
):
# Session exists and is not closing: create request from session
session = self.session_controller.get(session_id) session = self.session_controller.get(session_id)
req = session.create_req( req = session.create_req(
recv_req, recv_req,
@@ -1866,7 +1869,13 @@ class Scheduler(
return return
else: else:
# Session ID provided but session not found # Session not found, or session is closing
if session_id in self.session_controller:
error_msg = (
f"Invalid request: close was requested for session {session_id}"
)
else:
error_msg = f"Invalid request: session id {session_id} does not exist"
req = Req( req = Req(
recv_req.rid, recv_req.rid,
recv_req.input_text, recv_req.input_text,
@@ -1875,9 +1884,7 @@ class Scheduler(
vocab_size=self.model_config.vocab_size, vocab_size=self.model_config.vocab_size,
) )
req.tokenizer = self.tokenizer req.tokenizer = self.tokenizer
req.set_finish_with_abort( req.set_finish_with_abort(error_msg)
f"Invalid request: session id {session_id} does not exist"
)
self.init_req_max_new_tokens(req) self.init_req_max_new_tokens(req)
self._add_request_to_queue(req) self._add_request_to_queue(req)
return return
@@ -3461,7 +3468,10 @@ class Scheduler(
return ExpertDistributionReqOutput() return ExpertDistributionReqOutput()
def open_session(self, recv_req: OpenSessionReqInput): def open_session(self, recv_req: OpenSessionReqInput):
return self.session_controller.open(recv_req) output = self.session_controller.open(recv_req)
if self.pp_rank == 0 and self.tp_rank == 0 and self.attn_cp_rank == 0:
return output
return None
def close_session(self, recv_req: CloseSessionReqInput): def close_session(self, recv_req: CloseSessionReqInput):
self.session_controller.close(recv_req) self.session_controller.close(recv_req)
@@ -117,6 +117,13 @@ class PoolStats:
class SchedulerRuntimeCheckerMixin: class SchedulerRuntimeCheckerMixin:
def _alive_streaming_session_count(self: Scheduler) -> int:
return sum(
1
for session in self.session_controller.sessions.values()
if session.streaming
)
def _session_held_tokens(self: Scheduler) -> int: def _session_held_tokens(self: Scheduler) -> int:
if isinstance(self.tree_cache, SessionAwareCache): if isinstance(self.tree_cache, SessionAwareCache):
return self.tree_cache.session_held_tokens() return self.tree_cache.session_held_tokens()
@@ -451,6 +458,8 @@ class SchedulerRuntimeCheckerMixin:
return return
self.get_pool_stats().update_scheduler_stats(self.stats) self.get_pool_stats().update_scheduler_stats(self.stats)
self.stats.num_streaming_sessions = self._alive_streaming_session_count()
self.stats.streaming_session_held_tokens = self._session_held_tokens()
priority_enabled = self.enable_priority_scheduling priority_enabled = self.enable_priority_scheduling
self.stats.num_running_reqs = QueueCount.from_reqs( self.stats.num_running_reqs = QueueCount.from_reqs(
@@ -25,6 +25,7 @@ from sglang.srt.managers.io_struct import (
) )
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 from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.utils.common import log_info_on_rank0
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
@@ -92,6 +93,7 @@ class Session:
self.timeout = timeout self.timeout = timeout
self.last_active_time: float = time.monotonic() self.last_active_time: float = time.monotonic()
self.req_nodes: Dict[str, SessionReqNode] = {} self.req_nodes: Dict[str, SessionReqNode] = {}
self.close_on_finish: bool = False
def is_timed_out(self) -> bool: def is_timed_out(self) -> bool:
if self.timeout is None: if self.timeout is None:
@@ -275,6 +277,9 @@ class SessionController:
streaming=bool(recv_req.streaming), streaming=bool(recv_req.streaming),
timeout=recv_req.timeout, timeout=recv_req.timeout,
) )
log_info_on_rank0(
logger, f"Session opened: {session_id} (active={len(self.sessions)})"
)
return OpenSessionReqOutput(session_id, True) return OpenSessionReqOutput(session_id, True)
def close(self, recv_req: CloseSessionReqInput): def close(self, recv_req: CloseSessionReqInput):
@@ -286,10 +291,30 @@ class SessionController:
def _close(self, session_id: str): def _close(self, session_id: str):
session = self.sessions[session_id] session = self.sessions[session_id]
req = None
has_unfinished_request = False
if session.streaming and session.req_nodes: if session.streaming and session.req_nodes:
assert len(session.req_nodes) == 1 assert len(session.req_nodes) == 1
req = next(iter(session.req_nodes.values())).req req = next(iter(session.req_nodes.values())).req
if not req.finished(): if not req.finished():
has_unfinished_request = True
if has_unfinished_request:
# An in-flight request is still decoding on this session's KV
# memory. Freeing now would corrupt the scheduler. Mark the
# session for deferred cleanup: the request keeps its session
# reference so cache_finished_req takes the streaming path,
# and we schedule release_session for after it completes.
session.close_on_finish = True
logger.info(
"Deferring session close for %s (unfinished request)",
session_id,
)
return
# No active request -- safe to release immediately.
if session.streaming and session.req_nodes:
req = next(iter(session.req_nodes.values())).req
req.session = None req.session = None
# Release multimodal features held by session requests. # Release multimodal features held by session requests.
@@ -304,20 +329,46 @@ class SessionController:
node.req.multimodal_inputs = None node.req.multimodal_inputs = None
if isinstance(self.tree_cache, SessionAwareCache): if isinstance(self.tree_cache, SessionAwareCache):
self.tree_cache.release_session(session_id) self.tree_cache.release_session(
session_id, req if session.streaming else None
)
del self.sessions[session_id] del self.sessions[session_id]
log_info_on_rank0(
logger, f"Session closed: {session_id} (active={len(self.sessions)})"
)
def maybe_reap(self, now: float, interval: float = 1.0): def maybe_reap(self, now: float, interval: float = 1.0):
# reap sessions every second # reap sessions every second
if now - self._last_reap_time > interval: if now - self._last_reap_time > interval:
self._last_reap_time = now self._last_reap_time = now
# Finish deferred closes for sessions whose requests completed.
pending = [
sid
for sid, session in self.sessions.items()
if session.close_on_finish and self._all_requests_finished(session)
]
for sid in pending:
log_info_on_rank0(
logger, f"Deferred close ready for session {sid}, releasing."
)
# Reset close_on_finish so _close proceeds with the release.
self.sessions[sid].close_on_finish = False
self._close(sid)
timed_out = [ timed_out = [
sid for sid, session in self.sessions.items() if session.is_timed_out() sid for sid, session in self.sessions.items() if session.is_timed_out()
] ]
for sid in timed_out: for sid in timed_out:
logger.info(f"Session {sid} timed out, closing.") log_info_on_rank0(logger, f"Session {sid} timed out, closing.")
self._close(sid) self._close(sid)
@staticmethod
def _all_requests_finished(session: "Session") -> bool:
if not session.req_nodes:
return True
return all(node.req.finished() for node in session.req_nodes.values())
@staticmethod @staticmethod
def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs): def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs):
# For session requests, adjust mm_inputs offsets by the prefix length. # For session requests, adjust mm_inputs offsets by the prefix length.
@@ -1090,12 +1090,14 @@ class TokenizerCommunicatorMixin:
elif obj.session_id in self.session_futures: elif obj.session_id in self.session_futures:
return None return None
future = asyncio.Future()
self.session_futures[obj.session_id] = future
self.send_to_scheduler.send_pyobj(obj) self.send_to_scheduler.send_pyobj(obj)
self.session_futures[obj.session_id] = asyncio.Future() try:
session_id = await self.session_futures[obj.session_id] return await future
del self.session_futures[obj.session_id] finally:
return session_id self.session_futures.pop(obj.session_id, None)
async def close_session( async def close_session(
self: TokenizerManager, self: TokenizerManager,
@@ -2365,9 +2365,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
self.send_to_scheduler.send_pyobj(ranks) self.send_to_scheduler.send_pyobj(ranks)
def _handle_open_session_req_output(self, recv_obj): def _handle_open_session_req_output(self, recv_obj):
self.session_futures[recv_obj.session_id].set_result( future = self.session_futures.get(recv_obj.session_id)
recv_obj.session_id if recv_obj.success else None if future is None:
logger.warning(
"Open session response arrived after waiter cleanup: %s",
recv_obj.session_id,
) )
return
if not future.done():
future.set_result(recv_obj.session_id if recv_obj.success else None)
def _handle_update_weights_from_disk_req_output(self, recv_obj): def _handle_update_weights_from_disk_req_output(self, recv_obj):
if self.server_args.dp_size == 1: if self.server_args.dp_size == 1:
+44 -4
View File
@@ -9,6 +9,7 @@ import triton.language as tl
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import support_triton from sglang.srt.utils import support_triton
@@ -476,13 +477,52 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
req.mamba_pool_idx = None req.mamba_pool_idx = None
return return
# Streaming sessions transfer req_pool ownership into SessionSlot objects.
# Trim any speculative tail before that transfer, otherwise later turns
# restore only the committed prefix and can strand unreachable KV pages.
#
# Aborted streaming-session requests (e.g. input too long) skip the
# streaming path entirely. match_prefix did not restore the slot's KV
# state, so the request has a fresh pool slot that should be freed by
# cache_finished_req below (which also sets req_pool_idx = None).
from sglang.srt.managers.schedule_batch import FINISH_ABORT
is_streaming_session = (
isinstance(tree_cache, SessionAwareCache)
and getattr(req, "session", None) is not None
and req.session.streaming
)
is_aborted_streaming = is_streaming_session and isinstance(
getattr(req, "finished_reason", None), FINISH_ABORT
)
if is_streaming_session and not is_aborted_streaming:
start_p, end_p = req.pop_overallocated_kv_cache()
page_size = get_global_server_args().page_size
if page_size > 1:
start_p = ceil_align(start_p, page_size)
if start_p < end_p:
indices_to_free = tree_cache.req_to_token_pool.req_to_token[
req.req_pool_idx
][start_p:end_p]
tree_cache.token_to_kv_pool_allocator.free(indices_to_free)
req.kv_allocated_len = req.kv_committed_len
tree_cache.cache_finished_req(req, is_insert=is_insert) tree_cache.cache_finished_req(req, is_insert=is_insert)
# FIXME: SessionAwareCache.cache_finished_req sets req_pool_idx = None to # SessionAwareCache.cache_finished_req sets req_pool_idx = None to transfer
# transfer KV ownership to the SessionSlot, so we skip the remaining # KV ownership to the SessionSlot, so the remaining cleanup is skipped.
# cleanup (overalloc free + pool slot free). This means over-allocated # Streaming-session specific overalloc trimming must therefore happen
# tokens from speculative decoding are NOT freed between turns. # before cache_finished_req above.
if req.req_pool_idx is None: if req.req_pool_idx is None:
if is_streaming_session:
# The request no longer owns any KV once SessionAwareCache either
# transfers it into the session slot or frees it on abort. Mark
# both bookkeeping flags so busy-time memory checks do not keep
# counting this finished request as uncached KV.
if not req.kv_committed_freed:
req.pop_committed_kv_cache()
if not req.kv_overallocated_freed:
req.pop_overallocated_kv_cache()
return return
start_p, end_p = req.pop_overallocated_kv_cache() start_p, end_p = req.pop_overallocated_kv_cache()
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, Optional from typing import TYPE_CHECKING, Any, Dict, Optional
@@ -22,6 +23,9 @@ if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
logger = logging.getLogger(__name__)
class _VirtualNode: class _VirtualNode:
"""Sentinel node for streaming session requests. """Sentinel node for streaming session requests.
@@ -187,6 +191,16 @@ class SessionAwareCache(BasePrefixCache):
if slot is None or slot.req_pool_idx is None: if slot is None or slot.req_pool_idx is None:
return self.inner.match_prefix(params) return self.inner.match_prefix(params)
# If the request is destined for abort (e.g. input too long),
# do NOT restore the slot's KV state. set_finish_with_abort
# truncates origin_input_ids to [0], so alloc_for_extend would
# overwrite the slot's req_to_token row with a 1-token prefix,
# destroying the session's accumulated KV mapping. By skipping
# restore, the request gets a fresh pool slot from alloc_for_extend
# and the session slot remains untouched.
if req.to_finish is not None:
return self.inner.match_prefix(params)
slot.restore_to_req(req) slot.restore_to_req(req)
# logprob_start_len is already forced to -1 for streaming sessions # logprob_start_len is already forced to -1 for streaming sessions
@@ -208,13 +222,56 @@ class SessionAwareCache(BasePrefixCache):
if not _is_streaming(req): if not _is_streaming(req):
return self.inner.cache_finished_req(req, is_insert=is_insert, **kwargs) return self.inner.cache_finished_req(req, is_insert=is_insert, **kwargs)
from sglang.srt.managers.schedule_batch import FINISH_ABORT
session_id = req.session.session_id session_id = req.session.session_id
slot = self.slots.get(session_id) slot = self.slots.get(session_id)
is_first = slot is None is_first = slot is None
# When an aborted streaming-session request was scheduled (e.g.
# input too long), match_prefix skipped restore_to_req so the
# request got a fresh pool slot from alloc_for_extend. Don't
# overwrite the session slot -- free the transient KV and pool slot.
if not is_first and isinstance(req.finished_reason, FINISH_ABORT):
if req.req_pool_idx is not None:
# Free all KV pages allocated for this aborted request.
end = req.kv_allocated_len
if end > 0:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :end
]
self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free_slots.append(req.req_pool_idx)
req.req_pool_idx = None
return
if is_first: if is_first:
slot = SessionSlot() slot = SessionSlot()
self.slots[session_id] = slot self.slots[session_id] = slot
# If the session's KV is shrinking (e.g. client sent a shorter
# prompt after an abort), free the orphaned tail pages before
# save_from_req overwrites the slot's committed length.
# Never free tree-protected tokens — those are managed by the tree.
if (
not is_first
and slot.is_holding_kv
and req.kv_committed_len < slot.kv_committed_len
):
old_end = slot.kv_allocated_len
new_end = req.kv_committed_len
if self.page_size > 1:
new_end = ceil_align(new_end, self.page_size)
new_end = max(new_end, slot.cache_protected_len)
if new_end < old_end:
kv_indices = self.req_to_token_pool.req_to_token[
slot.req_pool_idx, new_end:old_end
]
self.token_to_kv_pool_allocator.free(kv_indices)
slot.cache_protected_len = min(
slot.cache_protected_len, req.kv_committed_len
)
slot.save_from_req(req, is_first=is_first) slot.save_from_req(req, is_first=is_first)
def cache_unfinished_req(self, req: Req, **kwargs): def cache_unfinished_req(self, req: Req, **kwargs):
@@ -251,23 +308,99 @@ class SessionAwareCache(BasePrefixCache):
# -- Session lifecycle -- # -- Session lifecycle --
def release_session(self, session_id: str): def _resolve_release_state(
self, slot: SessionSlot, req: Optional[Req]
) -> tuple[int, Any]:
"""Resolve the currently tree-owned prefix for a session slot.
A long-lived session can outlive radix-tree splits caused by unrelated
traffic. In that case, the saved `last_node` may no longer represent the
full protected prefix even though the slot's req_to_token row still
contains tree-owned indices at the front. Re-match the current request
text, then intersect the returned tree indices with the slot's row so
release uses the prefix that is still actually backed by the tree.
"""
protected_len = slot.cache_protected_len
lock_node = slot.last_node
# TODO: re-match logic disabled — match_prefix has side effects
# (splits) that disturb tree accounting. Directly using
# slot.last_node + cache_protected_len is safe after split analysis.
return protected_len, lock_node
if (
req is None
or not slot.is_holding_kv
or slot.req_pool_idx is None
or protected_len <= 0
):
return protected_len, lock_node
from sglang.srt.mem_cache.radix_cache import RadixKey
token_ids = (req.origin_input_ids + req.output_ids)[: slot.kv_committed_len]
if not token_ids:
return 0, None
match = self.inner.match_prefix(
MatchPrefixParams(
key=RadixKey(token_ids=token_ids, extra_key=req.extra_key),
req=None,
)
)
if len(match.device_indices) == 0:
return 0, None
max_protected_len = min(len(match.device_indices), protected_len)
row_indices = self.req_to_token_pool.req_to_token[
slot.req_pool_idx, :max_protected_len
].to(dtype=torch.int64)
match_indices = match.device_indices[:max_protected_len]
mismatches = (match_indices != row_indices).nonzero(as_tuple=False)
if mismatches.numel() == 0 and max_protected_len == len(match.device_indices):
common_len = max_protected_len
return common_len, match.last_device_node
common_len = (
int(mismatches[0].item()) if mismatches.numel() > 0 else max_protected_len
)
if self.page_size > 1:
common_len = (common_len // self.page_size) * self.page_size
if common_len <= 0:
return 0, None
rematch = self.inner.match_prefix(
MatchPrefixParams(
key=RadixKey(token_ids=token_ids[:common_len], extra_key=req.extra_key),
req=None,
)
)
return len(rematch.device_indices), rematch.last_device_node
def release_session(self, session_id: str, req: Optional[Req] = None):
"""Release all KV resources held by a streaming session.""" """Release all KV resources held by a streaming session."""
slot = self.slots.pop(session_id, None) slot = self.slots.pop(session_id, None)
if slot is None: if slot is None:
return return
protected_len, lock_node = self._resolve_release_state(slot, req)
tokens_freed = (
max(0, slot.kv_allocated_len - protected_len) if slot.is_holding_kv else 0
)
logger.info(
"Session KV released: %s (%d tokens freed)", session_id, tokens_freed
)
if slot.last_node is not None: if lock_node is not None:
if slot.swa_uuid_for_lock is not None: if slot.swa_uuid_for_lock is not None:
self.inner.dec_lock_ref( self.inner.dec_lock_ref(
slot.last_node, lock_node,
DecLockRefParams(swa_uuid_for_lock=slot.swa_uuid_for_lock), DecLockRefParams(swa_uuid_for_lock=slot.swa_uuid_for_lock),
) )
else: else:
self.inner.dec_lock_ref(slot.last_node) self.inner.dec_lock_ref(lock_node)
if slot.is_holding_kv: if slot.is_holding_kv:
start = slot.cache_protected_len start = protected_len
end = slot.kv_allocated_len end = slot.kv_allocated_len
if start < end: if start < end:
kv_indices = self.req_to_token_pool.req_to_token[ kv_indices = self.req_to_token_pool.req_to_token[
@@ -132,6 +132,10 @@ class SchedulerStats:
hicache_host_used_tokens: int = 0 hicache_host_used_tokens: int = 0
hicache_host_total_tokens: int = 0 hicache_host_total_tokens: int = 0
# Streaming session metrics
num_streaming_sessions: int = 0
streaming_session_held_tokens: int = 0
# Routing key metrics # Routing key metrics
num_unique_running_routing_keys: int = 0 num_unique_running_routing_keys: int = 0
routing_key_running_req_counts: List[int] = field(default_factory=list) routing_key_running_req_counts: List[int] = field(default_factory=list)
@@ -176,6 +180,7 @@ class SchedulerMetricsCollector:
labels: Dict[str, str], labels: Dict[str, str],
enable_lora: bool = False, enable_lora: bool = False,
enable_hierarchical_cache: bool = False, enable_hierarchical_cache: bool = False,
enable_streaming_session: bool = False,
server_args: Optional["ServerArgs"] = None, server_args: Optional["ServerArgs"] = None,
) -> None: ) -> None:
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR` # We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
@@ -184,6 +189,7 @@ class SchedulerMetricsCollector:
self.labels = labels self.labels = labels
self.enable_lora = enable_lora self.enable_lora = enable_lora
self.enable_hierarchical_cache = enable_hierarchical_cache self.enable_hierarchical_cache = enable_hierarchical_cache
self.enable_streaming_session = enable_streaming_session
self.last_log_time = time.perf_counter() self.last_log_time = time.perf_counter()
self._known_priorities: Set[int] = set() self._known_priorities: Set[int] = set()
@@ -654,6 +660,21 @@ class SchedulerMetricsCollector:
multiprocess_mode="mostrecent", multiprocess_mode="mostrecent",
) )
# Streaming session metrics (only created when streaming sessions are enabled)
if self.enable_streaming_session:
self.num_streaming_sessions = Gauge(
name="sglang:num_streaming_sessions",
documentation="The number of active streaming sessions.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.streaming_session_held_tokens = Gauge(
name="sglang:streaming_session_held_tokens",
documentation="The number of KV tokens currently held by streaming session slots.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.num_unique_running_routing_keys = Gauge( self.num_unique_running_routing_keys = Gauge(
name="sglang:num_unique_running_routing_keys", name="sglang:num_unique_running_routing_keys",
documentation="Number of unique routing keys in running batch.", documentation="Number of unique routing keys in running batch.",
@@ -1049,6 +1070,13 @@ class SchedulerMetricsCollector:
self.hicache_host_total_tokens, stats.hicache_host_total_tokens self.hicache_host_total_tokens, stats.hicache_host_total_tokens
) )
# Streaming session metrics (only logged if streaming sessions are enabled)
if self.enable_streaming_session:
self._log_gauge(self.num_streaming_sessions, stats.num_streaming_sessions)
self._log_gauge(
self.streaming_session_held_tokens, stats.streaming_session_held_tokens
)
self._log_gauge( self._log_gauge(
self.num_unique_running_routing_keys, stats.num_unique_running_routing_keys self.num_unique_running_routing_keys, stats.num_unique_running_routing_keys
) )
@@ -148,6 +148,7 @@ class SchedulerMetricsMixin:
labels=labels, labels=labels,
enable_lora=self.enable_lora, enable_lora=self.enable_lora,
enable_hierarchical_cache=self.enable_hierarchical_cache, enable_hierarchical_cache=self.enable_hierarchical_cache,
enable_streaming_session=self.server_args.enable_streaming_session,
server_args=self.server_args, server_args=self.server_args,
) )
self.enable_mfu_metrics = bool(self.server_args.enable_mfu_metrics) self.enable_mfu_metrics = bool(self.server_args.enable_mfu_metrics)
@@ -602,6 +603,8 @@ class SchedulerMetricsMixin:
self.stats.cache_hit_rate = cache_hit_rate self.stats.cache_hit_rate = cache_hit_rate
self.stats.max_total_num_tokens = self.max_total_num_tokens self.stats.max_total_num_tokens = self.max_total_num_tokens
self.stats.num_streaming_sessions = self._alive_streaming_session_count()
self.stats.streaming_session_held_tokens = self._session_held_tokens()
# Speculative decoding # Speculative decoding
self.stats.spec_accept_rate = spec_accept_rate self.stats.spec_accept_rate = spec_accept_rate
@@ -10,8 +10,12 @@ Usage:
""" """
import asyncio import asyncio
import json
import os
import tempfile
import time import time
import unittest import unittest
from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
import aiohttp import aiohttp
@@ -66,6 +70,21 @@ LEAK_FILLER = (
"We promptly judged antique ivory buckles for the next prize. " "We promptly judged antique ivory buckles for the next prize. "
) * 20 ) * 20
# ---------------------------------------------------------------------------
# Abort-heavy chunked prefill leak repro constants
# ---------------------------------------------------------------------------
ABORT_REPRO_CONTEXT_LEN = 512
ABORT_REPRO_PAGE_SIZE = 16
ABORT_REPRO_GEN_LEN = 8
ABORT_REPRO_SESSIONS = 4
ABORT_REPRO_WARMUP_TURNS = 2
ABORT_REPRO_ROUNDS = 8
ABORT_REPRO_STREAM_TOKENS = 150
ABORT_REPRO_ABORT_TOKENS = 320
ABORT_REPRO_NON_STREAMING_TOKENS = 96
ABORT_REPRO_CHUNKED_PREFILL_SIZE = 128
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Logprob leak helpers # Logprob leak helpers
@@ -206,6 +225,198 @@ async def _leak_run_all(base_url: str, tokenizer: Any) -> None:
assert resp.status == 200 assert resp.status == 200
def _make_token_sized_ids(
tokenizer: Any, prefix: str, min_tokens: int, max_tokens: Optional[int] = None
) -> list[int]:
text = prefix
chunk = " pack quartz wizard sphinx zebra fox " * 16
token_ids = tokenizer.encode(text)
while len(token_ids) < min_tokens:
text += chunk
token_ids = tokenizer.encode(text)
if max_tokens is not None:
token_ids = token_ids[:max_tokens]
return token_ids
async def _abort_repro_generate(
base_url: str,
session: aiohttp.ClientSession,
input_ids: list[int],
max_new_tokens: int,
session_params: Optional[dict[str, Any]] = None,
expect_abort: bool = False,
) -> Optional[dict[str, Any]]:
payload: dict[str, Any] = {
"input_ids": input_ids,
"sampling_params": {
"temperature": 0,
"max_new_tokens": max_new_tokens,
"no_stop_trim": True,
"skip_special_tokens": False,
},
}
if session_params:
payload["session_params"] = session_params
async with session.post(base_url + "/generate", json=payload) as resp:
text = await resp.text()
if expect_abort:
if resp.status == 200:
data = json.loads(text)
finish_reason = data.get("meta_info", {}).get("finish_reason", {})
assert finish_reason.get("type") == "abort", text
assert "maximum allowed length" in finish_reason.get(
"message", ""
), text
return data
assert resp.status == 400, text
assert "maximum allowed length" in text, text
return None
assert resp.status == 200, text
data = json.loads(text)
finish_reason = data.get("meta_info", {}).get("finish_reason", {})
assert finish_reason.get("type") != "abort", text
return data
def _read_tail(path: str, num_lines: int = 120) -> str:
if not path or not os.path.exists(path):
return "<missing log>"
lines = Path(path).read_text(errors="replace").splitlines()
return "\n".join(lines[-num_lines:])
async def _abort_repro_run_all(base_url: str, tokenizer: Any) -> None:
timeout = aiohttp.ClientTimeout(total=300)
async with aiohttp.ClientSession(timeout=timeout) as http:
session_ids = []
for _ in range(ABORT_REPRO_SESSIONS):
async with http.post(
base_url + "/open_session",
json={"capacity_of_str_len": 50000, "streaming": True},
) as resp:
assert resp.status == 200, await resp.text()
session_ids.append(await resp.json())
try:
for warmup_turn in range(ABORT_REPRO_WARMUP_TURNS):
warmup_tasks = []
for session_idx, session_id in enumerate(session_ids):
input_ids = _make_token_sized_ids(
tokenizer,
prefix=f"[warmup={warmup_turn} session={session_idx}]",
min_tokens=ABORT_REPRO_STREAM_TOKENS,
max_tokens=ABORT_REPRO_STREAM_TOKENS + 8,
)
warmup_tasks.append(
_abort_repro_generate(
base_url,
http,
input_ids,
ABORT_REPRO_GEN_LEN,
session_params={"id": session_id, "rid": None},
)
)
await asyncio.gather(*warmup_tasks)
for round_idx in range(ABORT_REPRO_ROUNDS):
mixed_tasks = []
for session_idx, session_id in enumerate(session_ids):
input_ids = _make_token_sized_ids(
tokenizer,
prefix=f"[round={round_idx} ok session={session_idx}]",
min_tokens=ABORT_REPRO_STREAM_TOKENS,
max_tokens=ABORT_REPRO_STREAM_TOKENS + 8,
)
mixed_tasks.append(
_abort_repro_generate(
base_url,
http,
input_ids,
ABORT_REPRO_GEN_LEN,
session_params={"id": session_id, "rid": None},
)
)
for ns_idx in range(2):
input_ids = _make_token_sized_ids(
tokenizer,
prefix=f"[round={round_idx} ns={ns_idx}]",
min_tokens=ABORT_REPRO_NON_STREAMING_TOKENS,
max_tokens=ABORT_REPRO_NON_STREAMING_TOKENS + 8,
)
mixed_tasks.append(
_abort_repro_generate(
base_url,
http,
input_ids,
ABORT_REPRO_GEN_LEN,
)
)
await asyncio.gather(*mixed_tasks)
abort_tasks = []
for session_idx, session_id in enumerate(session_ids):
input_ids = _make_token_sized_ids(
tokenizer,
prefix=f"[round={round_idx} abort session={session_idx}]",
min_tokens=ABORT_REPRO_ABORT_TOKENS,
)
abort_tasks.append(
_abort_repro_generate(
base_url,
http,
input_ids,
ABORT_REPRO_GEN_LEN,
session_params={"id": session_id, "rid": None},
expect_abort=True,
)
)
await asyncio.gather(*abort_tasks)
recovery_tasks = []
for session_idx, session_id in enumerate(session_ids):
input_ids = _make_token_sized_ids(
tokenizer,
prefix=f"[round={round_idx} recover session={session_idx}]",
min_tokens=ABORT_REPRO_NON_STREAMING_TOKENS,
max_tokens=ABORT_REPRO_NON_STREAMING_TOKENS + 8,
)
recovery_tasks.append(
_abort_repro_generate(
base_url,
http,
input_ids,
ABORT_REPRO_GEN_LEN,
session_params={"id": session_id, "rid": None},
)
)
recovery_results = await asyncio.gather(*recovery_tasks)
for result in recovery_results:
assert result is not None
assert result["meta_info"]["cached_tokens"] > 0, result
health = requests.get(base_url + "/health", timeout=10)
if health.status_code != 200:
raise RuntimeError(
f"server unhealthy after round={round_idx}: "
f"{health.status_code} {health.text}"
)
finally:
for session_id in session_ids:
async with http.post(
base_url + "/close_session", json={"session_id": session_id}
) as resp:
assert resp.status == 200, await resp.text()
# ===================================================================
# Test class
# ===================================================================
class TestStreamingSession(CustomTestCase): class TestStreamingSession(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -467,5 +678,94 @@ class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
kill_process_tree(cls.process.pid) kill_process_tree(cls.process.pid)
class TestStreamingSessionAbortLeakRepro(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
cls.stdout = tempfile.NamedTemporaryFile(
prefix="streaming-session-abort-repro.",
suffix=".stdout.log",
delete=False,
mode="w+",
encoding="utf-8",
)
cls.stderr = tempfile.NamedTemporaryFile(
prefix="streaming-session-abort-repro.",
suffix=".stderr.log",
delete=False,
mode="w+",
encoding="utf-8",
)
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--enable-streaming-session",
"--chunked-prefill-size",
str(ABORT_REPRO_CHUNKED_PREFILL_SIZE),
"--context-length",
str(ABORT_REPRO_CONTEXT_LEN),
"--page-size",
str(ABORT_REPRO_PAGE_SIZE),
"--max-running-requests",
"32",
"--log-level",
"info",
],
return_stdout_stderr=(cls.stdout, cls.stderr),
)
cls.tokenizer = get_tokenizer(cls.model)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
for handle in (cls.stdout, cls.stderr):
path = handle.name
handle.close()
if os.path.exists(path):
os.remove(path)
def test_abort_heavy_chunked_prefill_does_not_leak(self) -> None:
requests.post(self.base_url + "/flush_cache")
asyncio.run(_abort_repro_run_all(self.base_url, self.tokenizer))
for i in range(3):
ids = self.tokenizer.encode(f"Post-session cleanup request {i}.")
response = requests.post(
self.base_url + "/generate",
json={
"input_ids": ids,
"sampling_params": {"temperature": 0, "max_new_tokens": 4},
},
timeout=30,
)
self.assertEqual(response.status_code, 200, response.text)
time.sleep(5)
self.assertIsNone(
self.process.poll(),
"Server crashed during abort-heavy streaming session repro.\n"
f"---- stderr tail ----\n{_read_tail(self.stderr.name)}",
)
health = requests.get(self.base_url + "/health", timeout=10)
self.assertEqual(
health.status_code,
200,
"Server unhealthy after abort-heavy streaming session cleanup.\n"
f"---- stderr tail ----\n{_read_tail(self.stderr.name)}",
)
stderr_tail = _read_tail(self.stderr.name)
self.assertNotIn(
"token_to_kv_pool_allocator memory leak detected",
stderr_tail,
stderr_tail,
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -0,0 +1,358 @@
from types import SimpleNamespace
import torch
from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.mem_cache.base_prefix_cache import MatchResult
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache, SessionSlot
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=8, suite="stage-a-test-cpu")
class _FakeAllocator:
def __init__(self):
self.freed = []
def free(self, free_index: torch.Tensor):
self.freed.append(free_index.clone())
class _FakeInnerCache:
def __init__(self, req_to_token_pool, allocator, page_size, match_results=None):
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = allocator
self.page_size = page_size
self.match_results = list(match_results or [])
self.dec_lock_ref_calls = []
def cache_finished_req(self, *args, **kwargs):
raise AssertionError("Streaming requests should not delegate to inner cache")
def match_prefix(self, *args, **kwargs):
if not self.match_results:
raise AssertionError("Unexpected match_prefix call")
return self.match_results.pop(0)
def dec_lock_ref(self, node, *args, **kwargs):
self.dec_lock_ref_calls.append(node)
def supports_mamba(self):
return False
def sanity_check(self):
return None
class _FakeReq:
def __init__(
self, session_id: str, req_pool_idx: int, committed: int, allocated: int
):
self.session = SimpleNamespace(session_id=session_id, streaming=True)
self.req_pool_idx = req_pool_idx
self.kv_committed_len = committed
self.kv_allocated_len = allocated
self.kv_committed_freed = False
self.kv_overallocated_freed = False
self.origin_input_ids = list(range(committed))
self.output_ids = []
self.extra_key = None
self.swa_evicted_seqlen = 0
self.last_node = None
self.cache_protected_len = 0
self.swa_uuid_for_lock = None
self.mamba_pool_idx = None
self.mamba_ping_pong_track_buffer = None
self.mamba_next_track_idx = None
self.mamba_last_track_seqlen = None
self.mamba_branching_seqlen = None
self.pop_overallocated_calls = 0
self.to_finish = None
self.finished_reason = None
def pop_committed_kv_cache(self):
assert not self.kv_committed_freed
self.kv_committed_freed = True
return self.kv_committed_len
def pop_overallocated_kv_cache(self):
assert not self.kv_overallocated_freed
self.pop_overallocated_calls += 1
self.kv_overallocated_freed = True
return self.kv_committed_len, self.kv_allocated_len
def test_streaming_release_kv_cache_trims_overallocated_tail(monkeypatch):
page_size = 16
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
tree_cache = SessionAwareCache(
_FakeInnerCache(req_to_token_pool, allocator, page_size)
)
req = _FakeReq("session-a", req_pool_idx=0, committed=17, allocated=40)
monkeypatch.setattr(
"sglang.srt.mem_cache.common.get_global_server_args",
lambda: SimpleNamespace(page_size=page_size, speculative_algorithm="eagle"),
)
release_kv_cache(req, tree_cache)
slot = tree_cache.slots["session-a"]
assert req.pop_overallocated_calls == 1
assert req.kv_committed_freed is True
assert req.kv_overallocated_freed is True
assert req.req_pool_idx is None
assert slot.kv_committed_len == 17
assert slot.kv_allocated_len == 17
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(32, 40))
def test_release_session_recomputes_current_tree_owned_prefix():
page_size = 16
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
full_match = MatchResult(
device_indices=torch.tensor(list(range(16)) + list(range(64, 96))),
last_device_node="stale-expanded",
last_host_node="stale-expanded",
)
protected_match = MatchResult(
device_indices=torch.tensor(list(range(16))),
last_device_node="current-protected",
last_host_node="current-protected",
)
inner = _FakeInnerCache(
req_to_token_pool,
allocator,
page_size,
match_results=[full_match, protected_match],
)
tree_cache = SessionAwareCache(inner)
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=48,
kv_allocated_len=48,
last_node="outdated-node",
cache_protected_len=32,
)
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
tree_cache.release_session("session-a", req)
assert inner.dec_lock_ref_calls == ["current-protected"]
assert req_to_token_pool.free_slots == [0]
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(16, 48))
def test_release_session_never_grows_tree_owned_prefix():
page_size = 16
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
overmatched = MatchResult(
device_indices=torch.tensor(list(range(48))),
last_device_node="overmatched-node",
last_host_node="overmatched-node",
)
capped_match = MatchResult(
device_indices=torch.tensor(list(range(16))),
last_device_node="original-lock-node",
last_host_node="original-lock-node",
)
inner = _FakeInnerCache(
req_to_token_pool,
allocator,
page_size,
match_results=[overmatched, capped_match],
)
tree_cache = SessionAwareCache(inner)
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=48,
kv_allocated_len=48,
last_node="outdated-node",
cache_protected_len=16,
)
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
tree_cache.release_session("session-a", req)
assert inner.dec_lock_ref_calls == ["original-lock-node"]
assert req_to_token_pool.free_slots == [0]
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(16, 48))
def test_match_prefix_abort_does_not_restore_live_session_slot():
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
inner = _FakeInnerCache(
req_to_token_pool,
allocator,
page_size=16,
match_results=[
MatchResult(
device_indices=torch.tensor([], dtype=torch.int64),
last_device_node=None,
last_host_node=None,
)
],
)
tree_cache = SessionAwareCache(inner)
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=48,
kv_allocated_len=48,
cache_protected_len=16,
)
req = _FakeReq("session-a", req_pool_idx=1, committed=1, allocated=1)
req.to_finish = FINISH_ABORT("too long")
result = tree_cache.match_prefix(
SimpleNamespace(
req=req,
key=SimpleNamespace(token_ids=list(range(64))),
)
)
slot = tree_cache.slots["session-a"]
assert req.req_pool_idx == 1
assert req.kv_committed_len == 1
assert req.kv_allocated_len == 1
assert slot.req_pool_idx == 0
assert slot.kv_committed_len == 48
assert slot.kv_allocated_len == 48
assert len(result.device_indices) == 0
def test_aborted_streaming_turn_preserves_slot_and_accounting(monkeypatch):
page_size = 16
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
tree_cache = SessionAwareCache(
_FakeInnerCache(req_to_token_pool, allocator, page_size)
)
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=48,
kv_allocated_len=48,
cache_protected_len=16,
swa_evicted_seqlen=8,
last_node="lock-node",
)
req = _FakeReq("session-a", req_pool_idx=1, committed=5, allocated=23)
req.finished_reason = FINISH_ABORT("too long")
monkeypatch.setattr(
"sglang.srt.mem_cache.common.get_global_server_args",
lambda: SimpleNamespace(page_size=page_size, speculative_algorithm="eagle"),
)
release_kv_cache(req, tree_cache)
slot = tree_cache.slots["session-a"]
assert slot.req_pool_idx == 0
assert slot.kv_committed_len == 48
assert slot.kv_allocated_len == 48
assert req.kv_committed_freed is True
assert req.kv_overallocated_freed is True
assert req.req_pool_idx is None
assert req.pop_overallocated_calls == 1
assert tree_cache.session_held_tokens() == 32
assert tree_cache.session_held_full_tokens() == 32
assert tree_cache.session_held_swa_tokens() == 32
assert tree_cache.session_held_req_count() == 1
assert req_to_token_pool.free_slots == [1]
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(128, 151))
tree_cache.release_session("session-a")
assert tree_cache.session_held_tokens() == 0
assert tree_cache.session_held_swa_tokens() == 0
assert tree_cache.session_held_req_count() == 0
assert req_to_token_pool.free_slots == [1, 0]
assert len(allocator.freed) == 2
assert allocator.freed[1].tolist() == list(range(16, 48))
def test_session_shrink_frees_orphaned_tail():
"""When a session's KV shrinks (client retried with shorter prompt),
the orphaned tail pages must be freed before save_from_req overwrites
the slot."""
page_size = 16
pool_size = 256
req_to_token = torch.arange(pool_size, dtype=torch.int32).reshape(1, pool_size)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
inner = _FakeInnerCache(req_to_token_pool, allocator, page_size)
tree_cache = SessionAwareCache(inner)
# Session slot has 128 tokens committed
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=128,
kv_allocated_len=128,
last_node="lock-node",
cache_protected_len=16,
)
# New request finished with only 48 tokens (client truncated)
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
tree_cache.cache_finished_req(req)
slot = tree_cache.slots["session-a"]
# Slot should now reflect the shrunk state
assert slot.kv_committed_len == 48
assert slot.kv_allocated_len == 48
# The tail [48:128] should have been freed (page-aligned: [48:128])
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(48, 128))
def test_session_shrink_page_aligns_free_start():
"""The shrink free should page-align the start to avoid freeing
tokens that are still part of the new committed prefix."""
page_size = 16
pool_size = 256
req_to_token = torch.arange(pool_size, dtype=torch.int32).reshape(1, pool_size)
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
allocator = _FakeAllocator()
inner = _FakeInnerCache(req_to_token_pool, allocator, page_size)
tree_cache = SessionAwareCache(inner)
# Session slot has 128 tokens
tree_cache.slots["session-a"] = SessionSlot(
req_pool_idx=0,
kv_committed_len=128,
kv_allocated_len=128,
last_node="lock-node",
cache_protected_len=16,
)
# New request committed 50 tokens (not page-aligned)
req = _FakeReq("session-a", req_pool_idx=0, committed=50, allocated=50)
tree_cache.cache_finished_req(req)
slot = tree_cache.slots["session-a"]
assert slot.kv_committed_len == 50
# Free start should be ceil_align(50, 16) = 64, not 50
# So freed range is [64:128]
assert len(allocator.freed) == 1
assert allocator.freed[0].tolist() == list(range(64, 128))