Refactor streaming session abort handling (#22790)
This commit is contained in:
@@ -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())
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user