Fix: abort handling for dispatched requests after client disconnect (#35255)
Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
co-authored by
Xinyuan Tong
cctry
parent
e787de5478
commit
f478b2bb2d
@@ -3348,6 +3348,11 @@ class Scheduler(
|
|||||||
# it. Drop the marker once the request is actually gone.
|
# it. Drop the marker once the request is actually gone.
|
||||||
if req.finished() or not req.kv.holds_kv:
|
if req.finished() or not req.kv.holds_kv:
|
||||||
self._pending_chunked_abort_req = None
|
self._pending_chunked_abort_req = None
|
||||||
|
return
|
||||||
|
# The request moved to another scheduler queue after abort_request
|
||||||
|
# deferred it, so retry against its current location.
|
||||||
|
self._pending_chunked_abort_req = None
|
||||||
|
self.abort_request(AbortReq(rid=req.rid))
|
||||||
return
|
return
|
||||||
|
|
||||||
prepare_abort(req, "Aborted")
|
prepare_abort(req, "Aborted")
|
||||||
|
|||||||
@@ -233,6 +233,9 @@ class ReqState:
|
|||||||
last_completion_tokens: int = 1
|
last_completion_tokens: int = 1
|
||||||
ttft_observed: bool = False
|
ttft_observed: bool = False
|
||||||
|
|
||||||
|
dispatched: bool = False
|
||||||
|
abort_sent: bool = False
|
||||||
|
|
||||||
# For streaming output
|
# For streaming output
|
||||||
last_output_offset: int = 0
|
last_output_offset: int = 0
|
||||||
|
|
||||||
@@ -799,6 +802,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self._init_req_state(obj, request)
|
self._init_req_state(obj, request)
|
||||||
|
request_rids = {obj.rid} if obj.is_single else set(obj.rid)
|
||||||
try:
|
try:
|
||||||
if get_disagg().language_only:
|
if get_disagg().language_only:
|
||||||
self._handle_epd_disaggregation_encode_request(obj)
|
self._handle_epd_disaggregation_encode_request(obj)
|
||||||
@@ -822,17 +826,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
async for response in self._wait_one_response(obj, request):
|
async for response in self._wait_one_response(obj, request):
|
||||||
yield response
|
yield response
|
||||||
else:
|
else:
|
||||||
async for response in self._handle_batch_request(obj, request):
|
async for response in self._handle_batch_request(
|
||||||
|
obj, request, request_rids
|
||||||
|
):
|
||||||
yield response
|
yield response
|
||||||
except BaseException:
|
except BaseException:
|
||||||
# _init_req_state created a rid_to_state entry per (sub-)request up
|
# _init_req_state created a rid_to_state entry per (sub-)request up
|
||||||
# front. The normal remover is the scheduler-response path
|
# front. The normal remover is the scheduler-response path
|
||||||
# (_handle_batch_output), so a failure *before* a request reaches the
|
# (_handle_batch_output), so a failure *before* a request reaches the
|
||||||
# scheduler -- e.g. input-length validation rejecting an over-context
|
# scheduler -- e.g. input-length validation rejecting an over-context
|
||||||
# request -- would otherwise leak those entries forever. Drop any that
|
# request -- would otherwise leak those entries forever. Drop
|
||||||
# are still pending; entries already removed on the normal completion
|
# undelivered states, but abort dispatched requests for scheduler-side
|
||||||
# path are left untouched (pop is a no-op).
|
# cleanup.
|
||||||
self._discard_pending_req_states(obj)
|
self._release_req_states_on_failure(request_rids)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def _detect_input_format(
|
def _detect_input_format(
|
||||||
@@ -1571,6 +1577,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
time_stats = tokenized_obj.time_stats
|
time_stats = tokenized_obj.time_stats
|
||||||
tokenized_obj.wrap_pickle_fields()
|
tokenized_obj.wrap_pickle_fields()
|
||||||
self._dispatch_to_scheduler(tokenized_obj)
|
self._dispatch_to_scheduler(tokenized_obj)
|
||||||
|
self._mark_state_dispatched(tokenized_obj.rid)
|
||||||
dispatched = True
|
dispatched = True
|
||||||
tokenized_obj.time_stats = time_stats
|
tokenized_obj.time_stats = time_stats
|
||||||
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
||||||
@@ -1578,6 +1585,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
if not dispatched:
|
if not dispatched:
|
||||||
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
|
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
|
||||||
|
|
||||||
|
def _mark_state_dispatched(self, rid: str):
|
||||||
|
"""Record that *rid* reached the scheduler.
|
||||||
|
|
||||||
|
Only dispatched requests are aborted (not discarded) by the
|
||||||
|
handler-failure cleanup; see _release_req_states_on_failure.
|
||||||
|
"""
|
||||||
|
state = self.rid_to_state.get(rid)
|
||||||
|
if state is not None:
|
||||||
|
state.dispatched = True
|
||||||
|
|
||||||
async def _send_batch_request(
|
async def _send_batch_request(
|
||||||
self,
|
self,
|
||||||
tokenized_objs: List[
|
tokenized_objs: List[
|
||||||
@@ -1605,6 +1622,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
||||||
|
|
||||||
self._dispatch_to_scheduler(batch_req)
|
self._dispatch_to_scheduler(batch_req)
|
||||||
|
for tokenized_obj in tokenized_objs:
|
||||||
|
self._mark_state_dispatched(tokenized_obj.rid)
|
||||||
dispatched = True
|
dispatched = True
|
||||||
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
|
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
|
||||||
tokenized_obj.time_stats = time_stat
|
tokenized_obj.time_stats = time_stat
|
||||||
@@ -1819,7 +1838,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self,
|
self,
|
||||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||||
request: Optional[fastapi.Request] = None,
|
request: Optional[fastapi.Request] = None,
|
||||||
|
request_rids: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
|
if request_rids is None:
|
||||||
|
request_rids = set(obj.rid)
|
||||||
batch_size = obj.batch_size
|
batch_size = obj.batch_size
|
||||||
|
|
||||||
generators = []
|
generators = []
|
||||||
@@ -1885,6 +1907,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
tokenized_obj.sampling_params.max_new_tokens = 0
|
tokenized_obj.sampling_params.max_new_tokens = 0
|
||||||
tokenized_obj.stream = False
|
tokenized_obj.stream = False
|
||||||
self._init_req_state(tmp_obj)
|
self._init_req_state(tmp_obj)
|
||||||
|
request_rids.add(tmp_obj.rid)
|
||||||
await self._send_one_request(tokenized_obj)
|
await self._send_one_request(tokenized_obj)
|
||||||
await self._wait_one_response(tmp_obj, request).__anext__()
|
await self._wait_one_response(tmp_obj, request).__anext__()
|
||||||
|
|
||||||
@@ -1901,6 +1924,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
]
|
]
|
||||||
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
||||||
self._init_req_state(tmp_obj)
|
self._init_req_state(tmp_obj)
|
||||||
|
request_rids.add(tmp_obj.rid)
|
||||||
state = self.rid_to_state[tmp_obj.rid]
|
state = self.rid_to_state[tmp_obj.rid]
|
||||||
tokenized_obj.time_stats = state.time_stats
|
tokenized_obj.time_stats = state.time_stats
|
||||||
if tmp_obj.return_prompt_token_ids:
|
if tmp_obj.return_prompt_token_ids:
|
||||||
@@ -1970,14 +1994,21 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
if not abort_all and not rid:
|
if not abort_all and not rid:
|
||||||
logger.warning("Ignore abort_request with empty rid and abort_all=False")
|
logger.warning("Ignore abort_request with empty rid and abort_all=False")
|
||||||
return
|
return
|
||||||
if (
|
state = None if abort_all else self.rid_to_state.get(rid)
|
||||||
not abort_all
|
if not abort_all:
|
||||||
and get_serving().tokenizer_worker_num == 1
|
if state is not None:
|
||||||
and rid not in self.rid_to_state
|
if state.abort_sent:
|
||||||
):
|
return
|
||||||
return
|
state.abort_sent = True
|
||||||
|
elif get_serving().tokenizer_worker_num == 1:
|
||||||
|
return
|
||||||
req = AbortReq(rid=rid, abort_all=abort_all)
|
req = AbortReq(rid=rid, abort_all=abort_all)
|
||||||
self._dispatch_to_scheduler(req)
|
try:
|
||||||
|
self._dispatch_to_scheduler(req)
|
||||||
|
except BaseException:
|
||||||
|
if state is not None:
|
||||||
|
state.abort_sent = False
|
||||||
|
raise
|
||||||
if self.enable_metrics:
|
if self.enable_metrics:
|
||||||
# TODO: also use custom_labels from the request
|
# TODO: also use custom_labels from the request
|
||||||
self.metrics_collector.observe_one_aborted_request(
|
self.metrics_collector.observe_one_aborted_request(
|
||||||
@@ -2140,10 +2171,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Abort the request if the client is disconnected.
|
# Abort the request if the client is disconnected.
|
||||||
async def abort_request():
|
async def abort_request():
|
||||||
await asyncio.sleep(2)
|
await asyncio.sleep(2)
|
||||||
if obj.is_single:
|
rids = [obj.rid] if obj.is_single else obj.rid
|
||||||
self.abort_request(obj.rid)
|
for rid in rids:
|
||||||
else:
|
if rid in self.rid_to_state:
|
||||||
for rid in obj.rid:
|
|
||||||
self.abort_request(rid)
|
self.abort_request(rid)
|
||||||
|
|
||||||
background_tasks = BackgroundTasks()
|
background_tasks = BackgroundTasks()
|
||||||
@@ -3455,19 +3485,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header)
|
time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header)
|
||||||
time_stats.set_created_time(created_time)
|
time_stats.set_created_time(created_time)
|
||||||
|
|
||||||
def _discard_pending_req_states(self, obj):
|
def _release_req_states_on_failure(self, rids: Iterable[str]):
|
||||||
"""Drop rid_to_state entries created by _init_req_state for *obj*.
|
"""Release rid_to_state entries created for a failed handler.
|
||||||
|
|
||||||
Safe to call after a partial/failed dispatch: only entries still present
|
Undelivered states are removed locally. Dispatched requests are aborted
|
||||||
are removed, and the scheduler-response path looks up state with
|
and retained until the scheduler response removes them.
|
||||||
``.get(...)`` so a later output for a discarded rid is ignored, not fatal.
|
|
||||||
"""
|
"""
|
||||||
if not hasattr(obj, "is_single") or obj.is_single:
|
|
||||||
rids = [obj.rid]
|
|
||||||
else:
|
|
||||||
rids = obj.rid
|
|
||||||
for rid in rids:
|
for rid in rids:
|
||||||
self.rid_to_state.pop(rid, None)
|
state = self.rid_to_state.get(rid)
|
||||||
|
if state is None:
|
||||||
|
continue
|
||||||
|
if state.dispatched:
|
||||||
|
try:
|
||||||
|
self.abort_request(rid)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to abort request %s during cleanup", rid)
|
||||||
|
else:
|
||||||
|
del self.rid_to_state[rid]
|
||||||
|
|
||||||
def _should_dispatch_to_encoder(
|
def _should_dispatch_to_encoder(
|
||||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||||
|
|||||||
@@ -0,0 +1,69 @@
|
|||||||
|
"""Tests for deferred chunked-prefill aborts."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeReq:
|
||||||
|
"""Minimal stand-in for Req: only the fields the abort paths touch."""
|
||||||
|
|
||||||
|
def __init__(self, rid: str):
|
||||||
|
self.rid = rid
|
||||||
|
# Mirrors Req.kv; the abort paths read only these two predicates.
|
||||||
|
self.kv = SimpleNamespace(holds_kv=True, holds_mamba=False)
|
||||||
|
self.to_finish = None
|
||||||
|
self._finished = False
|
||||||
|
|
||||||
|
def finished(self):
|
||||||
|
return self._finished
|
||||||
|
|
||||||
|
|
||||||
|
def _make_scheduler(pending_req, *, chunked_req, running_reqs) -> Scheduler:
|
||||||
|
sched = Scheduler.__new__(Scheduler)
|
||||||
|
sched.chunked_req = chunked_req
|
||||||
|
sched._pending_chunked_abort_req = pending_req
|
||||||
|
sched.waiting_queue = []
|
||||||
|
sched.dllm_config = None
|
||||||
|
sched.grammar_manager = Mock()
|
||||||
|
sched.disaggregation_mode = None
|
||||||
|
sched.enable_hicache_storage = False
|
||||||
|
sched.mm_receiver = None
|
||||||
|
sched.ps = SimpleNamespace(pp_size=1)
|
||||||
|
sched.running_batch = SimpleNamespace(reqs=running_reqs)
|
||||||
|
sched.last_batch = None
|
||||||
|
return sched
|
||||||
|
|
||||||
|
|
||||||
|
class TestPendingChunkedAbortRace(CustomTestCase):
|
||||||
|
def test_req_left_chunked_slot_is_aborted(self):
|
||||||
|
req = _FakeReq("zombie_rid")
|
||||||
|
sched = _make_scheduler(req, chunked_req=None, running_reqs=[req])
|
||||||
|
|
||||||
|
sched.process_pending_chunked_abort()
|
||||||
|
|
||||||
|
self.assertIsNotNone(req.to_finish, "recorded abort was never applied")
|
||||||
|
self.assertIsNone(sched._pending_chunked_abort_req)
|
||||||
|
|
||||||
|
def test_finished_req_only_clears_marker(self):
|
||||||
|
req = _FakeReq("done_rid")
|
||||||
|
req._finished = True
|
||||||
|
sched = _make_scheduler(req, chunked_req=None, running_reqs=[])
|
||||||
|
|
||||||
|
sched.process_pending_chunked_abort()
|
||||||
|
|
||||||
|
self.assertIsNone(req.to_finish)
|
||||||
|
self.assertIsNone(sched._pending_chunked_abort_req)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main(verbosity=2)
|
||||||
@@ -10,11 +10,12 @@ Covers:
|
|||||||
- _handle_batch_output cleans up rid_to_state on finished requests
|
- _handle_batch_output cleans up rid_to_state on finished requests
|
||||||
- _init_req_state rejects duplicate rids
|
- _init_req_state rejects duplicate rids
|
||||||
- Resubmission succeeds after cleanup
|
- Resubmission succeeds after cleanup
|
||||||
|
- Handler failures clean up pending and dispatched requests
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import AsyncMock, MagicMock, Mock
|
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
@@ -454,41 +455,88 @@ def _make_generate_obj(rid, is_single):
|
|||||||
return obj
|
return obj
|
||||||
|
|
||||||
|
|
||||||
class TestDiscardPendingReqStates(CustomTestCase):
|
class TestReleaseReqStatesOnFailure(CustomTestCase):
|
||||||
"""Direct tests for _discard_pending_req_states."""
|
"""Direct tests for _release_req_states_on_failure."""
|
||||||
|
|
||||||
def test_discard_single(self):
|
def test_undelivered_single_is_dropped(self):
|
||||||
tm = _make_tokenizer_manager(self)
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "d_single"
|
rid = "d_single"
|
||||||
tm.rid_to_state[rid] = _make_req_state(rid)
|
tm.rid_to_state[rid] = _make_req_state(rid)
|
||||||
obj = Mock(spec=GenerateReqInput)
|
tm._release_req_states_on_failure([rid])
|
||||||
obj.is_single = True
|
|
||||||
obj.rid = rid
|
|
||||||
tm._discard_pending_req_states(obj)
|
|
||||||
self.assertNotIn(rid, tm.rid_to_state)
|
self.assertNotIn(rid, tm.rid_to_state)
|
||||||
|
|
||||||
def test_discard_batch_removes_all(self):
|
def test_undelivered_batch_removes_all(self):
|
||||||
tm = _make_tokenizer_manager(self)
|
tm = _make_tokenizer_manager(self)
|
||||||
rids = ["d0", "d1", "d2"]
|
rids = ["d0", "d1", "d2"]
|
||||||
for r in rids:
|
for r in rids:
|
||||||
tm.rid_to_state[r] = _make_req_state(r)
|
tm.rid_to_state[r] = _make_req_state(r)
|
||||||
obj = Mock(spec=GenerateReqInput)
|
tm._release_req_states_on_failure(rids)
|
||||||
obj.is_single = False
|
|
||||||
obj.rid = list(rids)
|
|
||||||
tm._discard_pending_req_states(obj)
|
|
||||||
for r in rids:
|
for r in rids:
|
||||||
self.assertNotIn(r, tm.rid_to_state)
|
self.assertNotIn(r, tm.rid_to_state)
|
||||||
|
|
||||||
def test_discard_ignores_already_removed(self):
|
def test_ignores_already_removed(self):
|
||||||
"""Popping a rid that is no longer present must not raise."""
|
"""A rid that is no longer present must not raise."""
|
||||||
tm = _make_tokenizer_manager(self)
|
tm = _make_tokenizer_manager(self)
|
||||||
tm.rid_to_state["p1"] = _make_req_state("p1")
|
tm.rid_to_state["p1"] = _make_req_state("p1")
|
||||||
obj = Mock(spec=GenerateReqInput)
|
tm._release_req_states_on_failure(["p1", "already_gone"])
|
||||||
obj.is_single = False
|
|
||||||
obj.rid = ["p1", "already_gone"]
|
|
||||||
tm._discard_pending_req_states(obj) # must not raise
|
|
||||||
self.assertNotIn("p1", tm.rid_to_state)
|
self.assertNotIn("p1", tm.rid_to_state)
|
||||||
|
|
||||||
|
def test_dispatched_single_is_aborted_and_state_kept(self):
|
||||||
|
tm = _make_tokenizer_manager(self)
|
||||||
|
tm.server_args.tokenizer_worker_num = 1
|
||||||
|
tm._dispatch_to_scheduler = Mock()
|
||||||
|
tm.enable_metrics = True
|
||||||
|
tm.metrics_collector = MagicMock()
|
||||||
|
rid = "d_live"
|
||||||
|
state = _make_req_state(rid)
|
||||||
|
state.dispatched = True
|
||||||
|
tm.rid_to_state[rid] = state
|
||||||
|
tm._release_req_states_on_failure([rid])
|
||||||
|
tm._release_req_states_on_failure([rid])
|
||||||
|
|
||||||
|
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
|
||||||
|
self.assertEqual(
|
||||||
|
[type(m) for m in sent], [AbortReq], "expected exactly one AbortReq"
|
||||||
|
)
|
||||||
|
self.assertEqual(sent[0].rid, rid)
|
||||||
|
self.assertIn(rid, tm.rid_to_state)
|
||||||
|
self.assertTrue(state.abort_sent)
|
||||||
|
tm.metrics_collector.observe_one_aborted_request.assert_called_once()
|
||||||
|
|
||||||
|
def test_dispatched_batch_aborts_delivered_and_drops_rest(self):
|
||||||
|
tm = _make_tokenizer_manager(self)
|
||||||
|
tm.server_args.tokenizer_worker_num = 1
|
||||||
|
tm._dispatch_to_scheduler = Mock()
|
||||||
|
delivered, undelivered = "d_delivered", "d_undelivered"
|
||||||
|
live = _make_req_state(delivered)
|
||||||
|
live.dispatched = True
|
||||||
|
tm.rid_to_state[delivered] = live
|
||||||
|
tm.rid_to_state[undelivered] = _make_req_state(undelivered)
|
||||||
|
tm._release_req_states_on_failure([delivered, undelivered])
|
||||||
|
|
||||||
|
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
|
||||||
|
self.assertEqual([type(m) for m in sent], [AbortReq])
|
||||||
|
self.assertEqual(sent[0].rid, delivered)
|
||||||
|
self.assertIn(delivered, tm.rid_to_state)
|
||||||
|
self.assertNotIn(undelivered, tm.rid_to_state)
|
||||||
|
|
||||||
|
def test_abort_failure_does_not_stop_cleanup(self):
|
||||||
|
tm = _make_tokenizer_manager(self)
|
||||||
|
tm.server_args.tokenizer_worker_num = 1
|
||||||
|
tm._dispatch_to_scheduler = Mock(side_effect=RuntimeError("send failed"))
|
||||||
|
delivered, undelivered = "live", "pending"
|
||||||
|
live = _make_req_state(delivered)
|
||||||
|
live.dispatched = True
|
||||||
|
tm.rid_to_state[delivered] = live
|
||||||
|
tm.rid_to_state[undelivered] = _make_req_state(undelivered)
|
||||||
|
|
||||||
|
with self.assertLogs(level="ERROR"):
|
||||||
|
tm._release_req_states_on_failure([delivered, undelivered])
|
||||||
|
|
||||||
|
self.assertIn(delivered, tm.rid_to_state)
|
||||||
|
self.assertFalse(live.abort_sent)
|
||||||
|
self.assertNotIn(undelivered, tm.rid_to_state)
|
||||||
|
|
||||||
|
|
||||||
class TestParallelStreamTaskCleanup(CustomTestCase):
|
class TestParallelStreamTaskCleanup(CustomTestCase):
|
||||||
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
|
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
|
||||||
@@ -595,6 +643,27 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
|
|||||||
for r in rids:
|
for r in rids:
|
||||||
self.assertNotIn(r, tm.rid_to_state)
|
self.assertNotIn(r, tm.rid_to_state)
|
||||||
|
|
||||||
|
def test_parallel_sampling_failure_cleans_generated_rid(self):
|
||||||
|
tm = _make_tm_for_generate(self)
|
||||||
|
obj = GenerateReqInput(
|
||||||
|
text=["hello"],
|
||||||
|
rid=["base"],
|
||||||
|
sampling_params={"n": 2},
|
||||||
|
)
|
||||||
|
tokenized = MagicMock()
|
||||||
|
tokenized.mm_inputs = None
|
||||||
|
tokenized.sampling_params = MagicMock()
|
||||||
|
tm._tokenize_one_request = AsyncMock(return_value=tokenized)
|
||||||
|
tm._send_one_request = Mock(side_effect=RuntimeError("dispatch failed"))
|
||||||
|
|
||||||
|
async def drive():
|
||||||
|
await tm.generate_request(obj).__anext__()
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "dispatch failed"):
|
||||||
|
asyncio.run(drive())
|
||||||
|
|
||||||
|
self.assertFalse(tm.rid_to_state)
|
||||||
|
|
||||||
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
|
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
|
||||||
tm = _make_tm_for_generate(self)
|
tm = _make_tm_for_generate(self)
|
||||||
obj = GenerateReqInput(
|
obj = GenerateReqInput(
|
||||||
@@ -641,5 +710,52 @@ class TestWaitOneResponseAfterStateFreed(CustomTestCase):
|
|||||||
self.assertEqual(out["text"], "hello")
|
self.assertEqual(out["text"], "hello")
|
||||||
|
|
||||||
|
|
||||||
|
class TestDisconnectAfterDispatchAbortsRequest(CustomTestCase):
|
||||||
|
"""Cancellation after dispatch must stop the scheduler request."""
|
||||||
|
|
||||||
|
@patch(
|
||||||
|
"sglang.srt.managers.tokenizer_manager.wrap_shm_features",
|
||||||
|
side_effect=lambda obj: obj,
|
||||||
|
)
|
||||||
|
def test_cancel_after_dispatch_sends_abort_and_keeps_state(self, _wrap_shm):
|
||||||
|
tm = _make_tm_for_generate(self)
|
||||||
|
tm.cuda_vmm_feature_transport = Mock()
|
||||||
|
tm.cuda_vmm_feature_transport.prepare_for_dispatch_async = AsyncMock(
|
||||||
|
return_value=[]
|
||||||
|
)
|
||||||
|
tm._dispatch_to_scheduler = Mock()
|
||||||
|
rid = "disconnect_zombie"
|
||||||
|
obj = _make_generate_obj(rid, is_single=True)
|
||||||
|
obj.return_prompt_token_ids = False
|
||||||
|
tokenized = MagicMock()
|
||||||
|
tokenized.rid = rid
|
||||||
|
tokenized.mm_inputs = None
|
||||||
|
tm._tokenize_one_request = AsyncMock(return_value=tokenized)
|
||||||
|
|
||||||
|
async def drive():
|
||||||
|
task = asyncio.create_task(tm.generate_request(obj).__anext__())
|
||||||
|
for _ in range(100):
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
if tm._dispatch_to_scheduler.called:
|
||||||
|
break
|
||||||
|
self.assertTrue(
|
||||||
|
tm._dispatch_to_scheduler.called, "request never dispatched"
|
||||||
|
)
|
||||||
|
state = tm.rid_to_state.get(rid)
|
||||||
|
self.assertIsNotNone(state)
|
||||||
|
self.assertTrue(state.dispatched)
|
||||||
|
|
||||||
|
task.cancel()
|
||||||
|
with self.assertRaises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
asyncio.run(drive())
|
||||||
|
|
||||||
|
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
|
||||||
|
aborts = [m for m in sent if isinstance(m, AbortReq) and m.rid == rid]
|
||||||
|
self.assertTrue(aborts, "disconnect must send an AbortReq to the scheduler")
|
||||||
|
self.assertIn(rid, tm.rid_to_state)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main(verbosity=2)
|
unittest.main(verbosity=2)
|
||||||
|
|||||||
@@ -376,6 +376,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
|
||||||
manager = object.__new__(TokenizerManager)
|
manager = object.__new__(TokenizerManager)
|
||||||
|
manager.rid_to_state = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
|
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
|
||||||
manager.cuda_vmm_feature_transport = transport
|
manager.cuda_vmm_feature_transport = transport
|
||||||
@@ -403,6 +404,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
|
manager.rid_to_state = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
manager._dispatch_to_scheduler = MagicMock(
|
manager._dispatch_to_scheduler = MagicMock(
|
||||||
side_effect=RuntimeError("send failed")
|
side_effect=RuntimeError("send failed")
|
||||||
@@ -437,6 +439,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
|
manager.rid_to_state = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
manager._dispatch_to_scheduler = MagicMock()
|
manager._dispatch_to_scheduler = MagicMock()
|
||||||
time_stats = MagicMock()
|
time_stats = MagicMock()
|
||||||
|
|||||||
@@ -78,9 +78,25 @@ def _returned_field_names(function):
|
|||||||
returned literal, assignments (annotated or not) to a returned name, a
|
returned literal, assignments (annotated or not) to a returned name, a
|
||||||
literal-key subscript write on it, and `.update(field=...)` on it. A
|
literal-key subscript write on it, and `.update(field=...)` on it. A
|
||||||
spelling this cannot see raises instead of skipping.
|
spelling this cannot see raises instead of skipping.
|
||||||
|
|
||||||
|
``overrides[name]`` is also accepted when ``name`` comes from
|
||||||
|
``for name in ("a", "b", ...)`` -- the keys stay statically enumerable.
|
||||||
"""
|
"""
|
||||||
names = set()
|
names = set()
|
||||||
returned = set()
|
returned = set()
|
||||||
|
# for x in ("a", "b"): ... -> {"x": {"a", "b"}}
|
||||||
|
loop_keys = {
|
||||||
|
node.target.id: {elt.value for elt in node.iter.elts}
|
||||||
|
for node in ast.walk(function)
|
||||||
|
if isinstance(node, ast.For)
|
||||||
|
and isinstance(node.target, ast.Name)
|
||||||
|
and isinstance(node.iter, (ast.Tuple, ast.List))
|
||||||
|
and node.iter.elts
|
||||||
|
and all(
|
||||||
|
isinstance(elt, ast.Constant) and isinstance(elt.value, str)
|
||||||
|
for elt in node.iter.elts
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
def top_level_keys(mapping):
|
def top_level_keys(mapping):
|
||||||
for key in mapping.keys:
|
for key in mapping.keys:
|
||||||
@@ -88,6 +104,14 @@ def _returned_field_names(function):
|
|||||||
raise AssertionError(f"non-literal key in {function.name}")
|
raise AssertionError(f"non-literal key in {function.name}")
|
||||||
names.add(key.value)
|
names.add(key.value)
|
||||||
|
|
||||||
|
def add_subscript_key(key):
|
||||||
|
if isinstance(key, ast.Constant):
|
||||||
|
names.add(key.value)
|
||||||
|
elif isinstance(key, ast.Name) and key.id in loop_keys:
|
||||||
|
names.update(loop_keys[key.id])
|
||||||
|
else:
|
||||||
|
raise AssertionError(f"non-literal key in {function.name}")
|
||||||
|
|
||||||
for node in ast.walk(function):
|
for node in ast.walk(function):
|
||||||
if isinstance(node, ast.Return) and node.value is not None:
|
if isinstance(node, ast.Return) and node.value is not None:
|
||||||
value = node.value
|
value = node.value
|
||||||
@@ -114,9 +138,7 @@ def _returned_field_names(function):
|
|||||||
if isinstance(target, ast.Subscript) and (
|
if isinstance(target, ast.Subscript) and (
|
||||||
isinstance(target.value, ast.Name) and target.value.id in returned
|
isinstance(target.value, ast.Name) and target.value.id in returned
|
||||||
):
|
):
|
||||||
if not isinstance(target.slice, ast.Constant):
|
add_subscript_key(target.slice)
|
||||||
raise AssertionError(f"non-literal key in {function.name}")
|
|
||||||
names.add(target.slice.value)
|
|
||||||
if (
|
if (
|
||||||
isinstance(node, ast.Call)
|
isinstance(node, ast.Call)
|
||||||
and isinstance(node.func, ast.Attribute)
|
and isinstance(node.func, ast.Attribute)
|
||||||
|
|||||||
Reference in New Issue
Block a user