diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b4118777e..1f9544800 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1848,8 +1848,11 @@ class Scheduler( self.stream_output([req], req.return_logprob) return - elif session_id in self.session_controller: - # Session exists: create request from session + elif ( + 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) req = session.create_req( recv_req, @@ -1866,7 +1869,13 @@ class Scheduler( return 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( recv_req.rid, recv_req.input_text, @@ -1875,9 +1884,7 @@ class Scheduler( 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" - ) + req.set_finish_with_abort(error_msg) self.init_req_max_new_tokens(req) self._add_request_to_queue(req) return @@ -3461,7 +3468,10 @@ class Scheduler( return ExpertDistributionReqOutput() 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): self.session_controller.close(recv_req) diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index 94aa0da06..189afae5b 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -117,6 +117,13 @@ class PoolStats: 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: if isinstance(self.tree_cache, SessionAwareCache): return self.tree_cache.session_held_tokens() @@ -451,6 +458,8 @@ class SchedulerRuntimeCheckerMixin: return 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 self.stats.num_running_reqs = QueueCount.from_reqs( diff --git a/python/sglang/srt/managers/session_controller.py b/python/sglang/srt/managers/session_controller.py index caf165c32..9bb763fce 100644 --- a/python/sglang/srt/managers/session_controller.py +++ b/python/sglang/srt/managers/session_controller.py @@ -25,6 +25,7 @@ from sglang.srt.managers.io_struct import ( ) from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache +from sglang.srt.utils.common import log_info_on_rank0 if TYPE_CHECKING: from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache @@ -92,6 +93,7 @@ class Session: self.timeout = timeout self.last_active_time: float = time.monotonic() self.req_nodes: Dict[str, SessionReqNode] = {} + self.close_on_finish: bool = False def is_timed_out(self) -> bool: if self.timeout is None: @@ -275,6 +277,9 @@ class SessionController: streaming=bool(recv_req.streaming), timeout=recv_req.timeout, ) + log_info_on_rank0( + logger, f"Session opened: {session_id} (active={len(self.sessions)})" + ) return OpenSessionReqOutput(session_id, True) def close(self, recv_req: CloseSessionReqInput): @@ -286,11 +291,31 @@ class SessionController: def _close(self, session_id: str): session = self.sessions[session_id] + req = None + has_unfinished_request = False 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 + 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 # Release multimodal features held by session requests. # Session reqs skip the normal mm cleanup path (scheduler and @@ -304,20 +329,46 @@ class SessionController: node.req.multimodal_inputs = None 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] + log_info_on_rank0( + logger, f"Session closed: {session_id} (active={len(self.sessions)})" + ) 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 + + # 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 = [ 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.") + log_info_on_rank0(logger, f"Session {sid} timed out, closing.") 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 def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs): # For session requests, adjust mm_inputs offsets by the prefix length. diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index b035cac37..8ec7ee4f6 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -1090,12 +1090,14 @@ class TokenizerCommunicatorMixin: elif obj.session_id in self.session_futures: return None + future = asyncio.Future() + self.session_futures[obj.session_id] = future self.send_to_scheduler.send_pyobj(obj) - self.session_futures[obj.session_id] = asyncio.Future() - session_id = await self.session_futures[obj.session_id] - del self.session_futures[obj.session_id] - return session_id + try: + return await future + finally: + self.session_futures.pop(obj.session_id, None) async def close_session( self: TokenizerManager, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 30e4ab9f9..bfc5bef63 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -2365,9 +2365,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): self.send_to_scheduler.send_pyobj(ranks) def _handle_open_session_req_output(self, recv_obj): - self.session_futures[recv_obj.session_id].set_result( - recv_obj.session_id if recv_obj.success else None - ) + future = self.session_futures.get(recv_obj.session_id) + 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): if self.server_args.dp_size == 1: diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 5a759ed11..391eea077 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -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.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.server_args import get_global_server_args 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 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) - # FIXME: SessionAwareCache.cache_finished_req sets req_pool_idx = None to - # transfer KV ownership to the SessionSlot, so we skip the remaining - # cleanup (overalloc free + pool slot free). This means over-allocated - # tokens from speculative decoding are NOT freed between turns. + # SessionAwareCache.cache_finished_req sets req_pool_idx = None to transfer + # KV ownership to the SessionSlot, so the remaining cleanup is skipped. + # Streaming-session specific overalloc trimming must therefore happen + # before cache_finished_req above. 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 start_p, end_p = req.pop_overallocated_kv_cache() diff --git a/python/sglang/srt/mem_cache/session_aware_cache.py b/python/sglang/srt/mem_cache/session_aware_cache.py index 983395ec6..cd5f01c70 100644 --- a/python/sglang/srt/mem_cache/session_aware_cache.py +++ b/python/sglang/srt/mem_cache/session_aware_cache.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Dict, Optional @@ -22,6 +23,9 @@ if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req +logger = logging.getLogger(__name__) + + class _VirtualNode: """Sentinel node for streaming session requests. @@ -187,6 +191,16 @@ class SessionAwareCache(BasePrefixCache): if slot is None or slot.req_pool_idx is None: 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) # logprob_start_len is already forced to -1 for streaming sessions @@ -208,13 +222,56 @@ class SessionAwareCache(BasePrefixCache): if not _is_streaming(req): 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 slot = self.slots.get(session_id) 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: slot = SessionSlot() 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) def cache_unfinished_req(self, req: Req, **kwargs): @@ -251,23 +308,99 @@ class SessionAwareCache(BasePrefixCache): # -- 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.""" slot = self.slots.pop(session_id, None) if slot is None: 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: self.inner.dec_lock_ref( - slot.last_node, + lock_node, DecLockRefParams(swa_uuid_for_lock=slot.swa_uuid_for_lock), ) else: - self.inner.dec_lock_ref(slot.last_node) + self.inner.dec_lock_ref(lock_node) if slot.is_holding_kv: - start = slot.cache_protected_len + start = protected_len end = slot.kv_allocated_len if start < end: kv_indices = self.req_to_token_pool.req_to_token[ diff --git a/python/sglang/srt/observability/metrics_collector.py b/python/sglang/srt/observability/metrics_collector.py index 9a81ee78f..3c3430216 100644 --- a/python/sglang/srt/observability/metrics_collector.py +++ b/python/sglang/srt/observability/metrics_collector.py @@ -132,6 +132,10 @@ class SchedulerStats: hicache_host_used_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 num_unique_running_routing_keys: int = 0 routing_key_running_req_counts: List[int] = field(default_factory=list) @@ -176,6 +180,7 @@ class SchedulerMetricsCollector: labels: Dict[str, str], enable_lora: bool = False, enable_hierarchical_cache: bool = False, + enable_streaming_session: bool = False, server_args: Optional["ServerArgs"] = None, ) -> None: # We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR` @@ -184,6 +189,7 @@ class SchedulerMetricsCollector: self.labels = labels self.enable_lora = enable_lora self.enable_hierarchical_cache = enable_hierarchical_cache + self.enable_streaming_session = enable_streaming_session self.last_log_time = time.perf_counter() self._known_priorities: Set[int] = set() @@ -654,6 +660,21 @@ class SchedulerMetricsCollector: 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( name="sglang:num_unique_running_routing_keys", 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 ) + # 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.num_unique_running_routing_keys, stats.num_unique_running_routing_keys ) diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index 593d6dd58..796c312c3 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -148,6 +148,7 @@ class SchedulerMetricsMixin: labels=labels, enable_lora=self.enable_lora, enable_hierarchical_cache=self.enable_hierarchical_cache, + enable_streaming_session=self.server_args.enable_streaming_session, server_args=self.server_args, ) 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.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 self.stats.spec_accept_rate = spec_accept_rate diff --git a/test/registered/sessions/test_streaming_session.py b/test/registered/sessions/test_streaming_session.py index 89c1a1fd8..59bb76981 100644 --- a/test/registered/sessions/test_streaming_session.py +++ b/test/registered/sessions/test_streaming_session.py @@ -10,8 +10,12 @@ Usage: """ import asyncio +import json +import os +import tempfile import time import unittest +from pathlib import Path from typing import Any, Optional import aiohttp @@ -66,6 +70,21 @@ LEAK_FILLER = ( "We promptly judged antique ivory buckles for the next prize. " ) * 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 @@ -206,6 +225,198 @@ async def _leak_run_all(base_url: str, tokenizer: Any) -> None: 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 "" + 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): @classmethod def setUpClass(cls): @@ -467,5 +678,94 @@ class TestStreamingSessionRetractMixedChunk(TestStreamingSession): 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__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_streaming_session_unit.py b/test/registered/unit/mem_cache/test_streaming_session_unit.py new file mode 100644 index 000000000..759cff862 --- /dev/null +++ b/test/registered/unit/mem_cache/test_streaming_session_unit.py @@ -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))