From aa78564e1a3ab684b7eafa3eba19a8654c90476f Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 15 Apr 2026 00:13:05 -0700 Subject: [PATCH] Refactor streaming session abort handling (#22790) --- .../sglang/srt/managers/session_controller.py | 37 ++- .../srt/mem_cache/session_aware_cache.py | 87 +++--- .../sessions/test_streaming_session.py | 282 +++++++++++++++++- .../mem_cache/test_streaming_session_unit.py | 193 +++++------- 4 files changed, 423 insertions(+), 176 deletions(-) diff --git a/python/sglang/srt/managers/session_controller.py b/python/sglang/srt/managers/session_controller.py index cd7e6f141..889e2d60b 100644 --- a/python/sglang/srt/managers/session_controller.py +++ b/python/sglang/srt/managers/session_controller.py @@ -94,6 +94,7 @@ class Session: self.last_active_time: float = time.monotonic() self.req_nodes: Dict[str, SessionReqNode] = {} self.close_on_finish: bool = False + self._inflight: bool = False def is_timed_out(self) -> bool: if self.timeout is None: @@ -117,7 +118,10 @@ class Session: abort_message = "" if self.streaming: # Streaming sessions: only simple appends allowed; reject otherwise. - if session_params.replace: + if self._inflight: + abort = True + abort_message = "Streaming session already has an active request." + elif session_params.replace: abort = True abort_message = "Streaming sessions do not support replace." elif session_params.drop_previous_output: @@ -130,7 +134,9 @@ class Session: abort_message = "Streaming sessions do not support offset." elif self.req_nodes: assert len(self.req_nodes) == 1 - _, last_req_node = self.req_nodes.popitem() + # Peek (don't pop) the single req_node. req_nodes is updated + # only in finish_req after the request completes successfully. + [last_req_node] = self.req_nodes.values() last_req = last_req_node.req elif session_params.replace: if session_params.rid is None: @@ -240,15 +246,27 @@ class Session: if abort: new_req.set_finish_with_abort(abort_message) elif self.streaming: - if last_req is not None: - last_req.session = None - self.req_nodes[req.rid] = SessionReqNode(new_req) + # req_nodes is NOT updated here — finish_req() handles it. + self._inflight = True else: new_req_node = SessionReqNode(new_req, last_req_node) self.req_nodes[req.rid] = new_req_node return new_req + def finish_req(self, req): + """Update req_nodes after a streaming request finishes successfully.""" + self._inflight = False + if self.req_nodes: + [prev_node] = self.req_nodes.values() + prev_node.req.session = None + self.req_nodes.clear() + self.req_nodes[req.rid] = SessionReqNode(req) + + def abort_req(self): + """Clear inflight flag on abort (req_nodes stays unchanged).""" + self._inflight = False + class SessionController: def __init__(self, tree_cache: BasePrefixCache): @@ -293,9 +311,12 @@ class SessionController: session = self.sessions[session_id] req = None has_unfinished_request = False - if session.streaming and session.req_nodes: + if session.streaming and session._inflight: + has_unfinished_request = True + elif session.streaming and session.req_nodes: assert len(session.req_nodes) == 1 - req = next(iter(session.req_nodes.values())).req + [last_node] = session.req_nodes.values() + req = last_node.req if not req.finished(): has_unfinished_request = True @@ -362,7 +383,7 @@ class SessionController: self._close(sid) @staticmethod - def _all_requests_finished(session: "Session") -> bool: + 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()) diff --git a/python/sglang/srt/mem_cache/session_aware_cache.py b/python/sglang/srt/mem_cache/session_aware_cache.py index 134bf99d1..36bde595b 100644 --- a/python/sglang/srt/mem_cache/session_aware_cache.py +++ b/python/sglang/srt/mem_cache/session_aware_cache.py @@ -183,14 +183,12 @@ 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. + # Pre-aborted req (scheduler-level abort, e.g. input too long): + # detach from session so cache_finished_req treats it as a normal + # req. The slot stays intact for the next request. if req.to_finish is not None: + req.session.abort_req() + req.session = None return self.inner.match_prefix(params) slot.restore_to_req(req) @@ -220,52 +218,51 @@ class SessionAwareCache(BasePrefixCache): 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 + # Mid-processing abort only. Pre-aborted reqs have session=None + # (set in match_prefix) and never reach here. + # Nuke all KV via release_session, delete slot. Token IDs stay + # in req_nodes (finish_req was never called -> last successful + # req). Next request re-prefills from scratch. + if isinstance(req.finished_reason, FINISH_ABORT): + if slot is None: + # First-request mid-processing abort: create ephemeral + # slot from req state so release_session handles cleanup. + # Include last_node/cache_protected_len from the req so + # release_session calls dec_lock_ref on the tree lock. + slot = SessionSlot( + req_pool_idx=req.req_pool_idx, + kv_allocated_len=req.kv_allocated_len, + last_node=req.last_node, + cache_protected_len=req.cache_protected_len, + swa_uuid_for_lock=req.swa_uuid_for_lock, + ) + self.slots[session_id] = slot + slot.kv_allocated_len = max(slot.kv_allocated_len, req.kv_allocated_len) + self.release_session(session_id) + req.req_pool_idx = None + req.session.abort_req() + self._mark_kv_freed(req) 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) + # Update req_nodes to this successfully finished request. + req.session.finish_req(req) + + self._mark_kv_freed(req) + + @staticmethod + def _mark_kv_freed(req: Req): + """Set bookkeeping flags so busy check skips this finished req.""" + if not req.kv_committed_freed: + req.pop_committed_kv_cache() + if not req.kv_overallocated_freed: + req.pop_overallocated_kv_cache() + def cache_unfinished_req(self, req: Req, **kwargs): if _is_streaming(req): # in chunked_prefill for streaming, we skip the stash path which triggers radix. diff --git a/test/registered/sessions/test_streaming_session.py b/test/registered/sessions/test_streaming_session.py index 8c8f82702..1d5a9fd62 100644 --- a/test/registered/sessions/test_streaming_session.py +++ b/test/registered/sessions/test_streaming_session.py @@ -73,13 +73,13 @@ LEAK_FILLER = ( ABORT_REPRO_CONTEXT_LEN = 512 ABORT_REPRO_PAGE_SIZE = 16 -ABORT_REPRO_GEN_LEN = 8 +ABORT_REPRO_GEN_LEN = 4 ABORT_REPRO_SESSIONS = 4 -ABORT_REPRO_WARMUP_TURNS = 2 +ABORT_REPRO_WARMUP_TURNS = 1 ABORT_REPRO_ROUNDS = 8 -ABORT_REPRO_STREAM_TOKENS = 150 -ABORT_REPRO_ABORT_TOKENS = 320 -ABORT_REPRO_NON_STREAMING_TOKENS = 96 +ABORT_REPRO_STREAM_TOKENS = 16 +ABORT_REPRO_ABORT_TOKENS = 600 +ABORT_REPRO_NON_STREAMING_TOKENS = 16 ABORT_REPRO_CHUNKED_PREFILL_SIZE = 128 @@ -265,10 +265,10 @@ async def _abort_repro_generate( assert finish_reason.get("type") == "abort", text assert "maximum allowed length" in finish_reason.get( "message", "" - ), text + ) or "context length" in finish_reason.get("message", ""), text return data assert resp.status == 400, text - assert "maximum allowed length" in text, text + assert "maximum allowed length" in text or "context length" in text, text return None assert resp.status == 200, text @@ -584,6 +584,274 @@ class TestStreamingSession(CustomTestCase): "likely a token memory leak from streaming session lifecycle.", ) + def test_nth_mid_abort_recovery(self) -> None: + """Abort a running streaming session request (nth turn) via the + abort API. Session rolls back to last successful turn.""" + requests.post(self.base_url + "/flush_cache") + + resp = requests.post( + self.base_url + "/open_session", + json={"capacity_of_str_len": 50000, "streaming": True}, + ) + self.assertEqual(resp.status_code, 200) + session_id = resp.json() + + try: + # Turn 1: normal generate to create slot. + ids_1 = self.tokenizer.encode("Tell me a very long story about a wizard.") + resp_1 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_1, + "sampling_params": {"temperature": 0, "max_new_tokens": 16}, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=30, + ) + self.assertEqual(resp_1.status_code, 200, resp_1.text) + data_1 = resp_1.json() + turn_1_total = ( + data_1["meta_info"]["prompt_tokens"] + + data_1["meta_info"]["completion_tokens"] + ) + + # Turn 2: long generate, then abort mid-decode. + ids_2 = self.tokenizer.encode(" Continue the story in great detail.") + + import threading + + result = [None] + + def do_generate(): + r = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_2, + "sampling_params": { + "temperature": 0, + "max_new_tokens": 100000, + }, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=60, + ) + result[0] = r + + t = threading.Thread(target=do_generate) + t.start() + time.sleep(0.5) + abort_resp = requests.post( + self.base_url + "/abort_request", + json={"rid": "", "abort_all": True}, + timeout=10, + ) + self.assertEqual(abort_resp.status_code, 200, abort_resp.text) + t.join(timeout=30) + + self.assertIsNotNone(result[0], "Turn 2 should have returned") + data_2 = result[0].json() + self.assertEqual( + data_2["meta_info"]["finish_reason"]["type"], + "abort", + "Turn 2 should be aborted, not finished normally", + ) + + # Turn 3: recovery. Rolls back to turn 1. + ids_3 = self.tokenizer.encode(" What happens next?") + for attempt in range(20): + resp_3 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_3, + "sampling_params": {"temperature": 0, "max_new_tokens": 8}, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=30, + ) + if resp_3.status_code == 200: + break + time.sleep(0.5) + self.assertEqual(resp_3.status_code, 200, resp_3.text) + data_3 = resp_3.json() + # prompt_tokens = turn_1_total + append (BOS stripped). + bos = 1 if ids_3[0] == self.tokenizer.bos_token_id else 0 + expected_prompt_3 = turn_1_total + len(ids_3) - bos + self.assertEqual( + data_3["meta_info"]["prompt_tokens"], + expected_prompt_3, + "prompt_tokens must equal turn_1_total + append (no stale abort context)", + ) + finally: + requests.post( + self.base_url + "/close_session", + json={"session_id": session_id}, + ) + + health = requests.get(self.base_url + "/health", timeout=10) + self.assertEqual(health.status_code, 200) + + def test_first_mid_abort_recovery(self) -> None: + """Abort the very first request on a streaming session mid-decode. + No slot exists yet (ephemeral slot created and nuked). + Verify the session is still usable afterward.""" + requests.post(self.base_url + "/flush_cache") + + resp = requests.post( + self.base_url + "/open_session", + json={"capacity_of_str_len": 50000, "streaming": True}, + ) + self.assertEqual(resp.status_code, 200) + session_id = resp.json() + + try: + ids_1 = self.tokenizer.encode("Tell me a very long story about a wizard.") + + import threading + + result = [None] + + def do_generate(): + r = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_1, + "sampling_params": { + "temperature": 0, + "max_new_tokens": 100000, + }, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=60, + ) + result[0] = r + + t = threading.Thread(target=do_generate) + t.start() + time.sleep(0.5) + abort_resp = requests.post( + self.base_url + "/abort_request", + json={"rid": "", "abort_all": True}, + timeout=10, + ) + self.assertEqual(abort_resp.status_code, 200, abort_resp.text) + t.join(timeout=30) + + self.assertIsNotNone(result[0], "Turn 1 should have returned") + data_1 = result[0].json() + self.assertEqual( + data_1["meta_info"]["finish_reason"]["type"], + "abort", + "Turn 1 should be aborted, not finished normally", + ) + + # Turn 2: recovery. No inherited context (req_nodes empty). + ids_2 = self.tokenizer.encode("Tell me a short joke.") + for attempt in range(20): + resp_2 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_2, + "sampling_params": {"temperature": 0, "max_new_tokens": 8}, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=30, + ) + if resp_2.status_code == 200: + break + time.sleep(0.5) + self.assertEqual(resp_2.status_code, 200, resp_2.text) + data_2 = resp_2.json() + self.assertEqual( + data_2["meta_info"]["prompt_tokens"], + len(ids_2), + "prompt_tokens must equal turn 2 input only (no inherited context)", + ) + finally: + requests.post( + self.base_url + "/close_session", + json={"session_id": session_id}, + ) + + health = requests.get(self.base_url + "/health", timeout=10) + self.assertEqual(health.status_code, 200) + + def test_preabort_recovery(self) -> None: + """Pre-aborted request (unsupported offset) does not corrupt session. + The slot is preserved, and the next turn inherits correctly.""" + requests.post(self.base_url + "/flush_cache") + + resp = requests.post( + self.base_url + "/open_session", + json={"capacity_of_str_len": 50000, "streaming": True}, + ) + self.assertEqual(resp.status_code, 200) + session_id = resp.json() + + try: + # Turn 1: normal generate to create slot. + ids_1 = self.tokenizer.encode("Tell me a very long story about a wizard.") + resp_1 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_1, + "sampling_params": {"temperature": 0, "max_new_tokens": 16}, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=30, + ) + self.assertEqual(resp_1.status_code, 200, resp_1.text) + data_1 = resp_1.json() + turn_1_total = ( + data_1["meta_info"]["prompt_tokens"] + + data_1["meta_info"]["completion_tokens"] + ) + + # Turn 2: pre-aborted via unsupported offset parameter. + ids_2 = self.tokenizer.encode(" This should be rejected.") + resp_2 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_2, + "sampling_params": {"temperature": 0, "max_new_tokens": 8}, + "session_params": { + "id": session_id, + "rid": None, + "offset": 1, + }, + }, + timeout=30, + ) + self.assertIn(resp_2.status_code, (200, 400), resp_2.text) + + # Turn 3: normal append. Slot should be intact from turn 1. + ids_3 = self.tokenizer.encode(" What happens next?") + resp_3 = requests.post( + self.base_url + "/generate", + json={ + "input_ids": ids_3, + "sampling_params": {"temperature": 0, "max_new_tokens": 8}, + "session_params": {"id": session_id, "rid": None}, + }, + timeout=30, + ) + self.assertEqual(resp_3.status_code, 200, resp_3.text) + data_3 = resp_3.json() + bos = 1 if ids_3[0] == self.tokenizer.bos_token_id else 0 + expected_prompt_3 = turn_1_total + len(ids_3) - bos + self.assertEqual( + data_3["meta_info"]["prompt_tokens"], + expected_prompt_3, + "prompt_tokens must equal turn_1_total + append (slot preserved)", + ) + finally: + requests.post( + self.base_url + "/close_session", + json={"session_id": session_id}, + ) + + health = requests.get(self.base_url + "/health", timeout=10) + self.assertEqual(health.status_code, 200) + class TestStreamingSessionMixedChunk(TestStreamingSession): """Streaming session with --enable-mixed-chunk.""" diff --git a/test/registered/unit/mem_cache/test_streaming_session_unit.py b/test/registered/unit/mem_cache/test_streaming_session_unit.py index ef0b786fc..2f32d5555 100644 --- a/test/registered/unit/mem_cache/test_streaming_session_unit.py +++ b/test/registered/unit/mem_cache/test_streaming_session_unit.py @@ -49,7 +49,13 @@ 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.session = SimpleNamespace( + session_id=session_id, + streaming=True, + finish_req=lambda req: None, + abort_req=lambda: None, + _inflight=False, + ) self.req_pool_idx = req_pool_idx self.kv_committed_len = committed self.kv_allocated_len = allocated @@ -83,7 +89,10 @@ class _FakeReq: return self.kv_committed_len, self.kv_allocated_len -def test_streaming_release_kv_cache_trims_overallocated_tail(monkeypatch): +def test_streaming_release_kv_cache_defers_tail_free(monkeypatch): + """Spec tail is NOT trimmed in cache_finished_req; it is deferred to + match_prefix's orphan tail free on the next turn. cache_finished_req + only sets bookkeeping flags and saves the slot as-is.""" 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=[]) @@ -101,17 +110,18 @@ def test_streaming_release_kv_cache_trims_overallocated_tail(monkeypatch): 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 + # Slot keeps the full allocation — tail free is deferred to match_prefix. 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)) + assert slot.kv_allocated_len == 40 + assert len(allocator.freed) == 0 -def test_match_prefix_abort_does_not_restore_live_session_slot(): +def test_preabort_detaches_session_and_preserves_slot(): + """Pre-aborted req (to_finish set before match_prefix) is detached from + the session: session=None, abort_req() called. Slot stays intact.""" 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() @@ -145,132 +155,83 @@ def test_match_prefix_abort_does_not_restore_live_session_slot(): ) ) + # Req detached from session. + assert req.session is None + # Slot untouched. 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 +def test_first_mid_abort_nukes_ephemeral_slot(): + """First-request mid-processing abort: no slot exists yet, ephemeral + slot is created from req state and nuked via release_session.""" + page_size = 1 + 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() + inner = _FakeInnerCache(req_to_token_pool, allocator, page_size) + tree_cache = SessionAwareCache(inner) + + # No slot exists yet (first request). + req = _FakeReq("session-a", req_pool_idx=0, committed=0, allocated=20) + req.finished_reason = FINISH_ABORT("input too long") + + tree_cache.cache_finished_req(req) + + # Slot must NOT be created. + assert "session-a" not in tree_cache.slots + # Transient pool slot freed. + assert req.req_pool_idx is None + assert req_to_token_pool.free_slots == [0] + assert len(allocator.freed) == 1 + assert allocator.freed[0].tolist() == list(range(20)) + # Bookkeeping flags set. + assert req.kv_committed_freed is True + assert req.kv_overallocated_freed is True + + +def test_nth_mid_abort_nukes_session_slot(): + """Nth-request mid-processing abort: slot exists, restore_to_req ran. + ALL KV is wiped (release_session). Slot is deleted. Token IDs stay + in req_nodes for next turn's re-prefill.""" + page_size = 1 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) - ) + inner = _FakeInnerCache(req_to_token_pool, allocator, page_size) + tree_cache = SessionAwareCache(inner) + + # Session already has a slot from a previous turn. 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", + kv_committed_len=50, + kv_allocated_len=50, + last_node=None, + cache_protected_len=0, ) - req = _FakeReq("session-a", req_pool_idx=1, committed=5, allocated=23) - req.finished_reason = FINISH_ABORT("too long") + # Mid-processing abort: req has the SESSION slot's pool_idx (restore_to_req ran). + req = _FakeReq("session-a", req_pool_idx=0, committed=60, allocated=65) + req.finished_reason = FINISH_ABORT("client disconnected") - monkeypatch.setattr( - "sglang.srt.mem_cache.common.get_global_server_args", - lambda: SimpleNamespace(page_size=page_size, speculative_algorithm="eagle"), - ) + tree_cache.cache_finished_req(req) - 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 + # Slot wiped — deleted from slots dict. + assert "session-a" not in tree_cache.slots + # All KV freed: [0, 65) from release_session (slot extended to req's allocated). + assert len(allocator.freed) == 1 + assert allocator.freed[0].tolist() == list(range(65)) + # Pool slot returned. + assert req_to_token_pool.free_slots == [0] + assert req.req_pool_idx is None + # Bookkeeping flags set. 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)) +# Shrink tests removed: streaming sessions are append-only after the +# rollback fix in session_controller (rollback_aborted_req). The shrink +# code path in cache_finished_req no longer exists.