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.
|
||||
if req.finished() or not req.kv.holds_kv:
|
||||
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
|
||||
|
||||
prepare_abort(req, "Aborted")
|
||||
|
||||
@@ -233,6 +233,9 @@ class ReqState:
|
||||
last_completion_tokens: int = 1
|
||||
ttft_observed: bool = False
|
||||
|
||||
dispatched: bool = False
|
||||
abort_sent: bool = False
|
||||
|
||||
# For streaming output
|
||||
last_output_offset: int = 0
|
||||
|
||||
@@ -799,6 +802,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
)
|
||||
|
||||
self._init_req_state(obj, request)
|
||||
request_rids = {obj.rid} if obj.is_single else set(obj.rid)
|
||||
try:
|
||||
if get_disagg().language_only:
|
||||
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):
|
||||
yield response
|
||||
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
|
||||
except BaseException:
|
||||
# _init_req_state created a rid_to_state entry per (sub-)request up
|
||||
# front. The normal remover is the scheduler-response path
|
||||
# (_handle_batch_output), so a failure *before* a request reaches the
|
||||
# scheduler -- e.g. input-length validation rejecting an over-context
|
||||
# request -- would otherwise leak those entries forever. Drop any that
|
||||
# are still pending; entries already removed on the normal completion
|
||||
# path are left untouched (pop is a no-op).
|
||||
self._discard_pending_req_states(obj)
|
||||
# request -- would otherwise leak those entries forever. Drop
|
||||
# undelivered states, but abort dispatched requests for scheduler-side
|
||||
# cleanup.
|
||||
self._release_req_states_on_failure(request_rids)
|
||||
raise
|
||||
|
||||
def _detect_input_format(
|
||||
@@ -1571,6 +1577,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
time_stats = tokenized_obj.time_stats
|
||||
tokenized_obj.wrap_pickle_fields()
|
||||
self._dispatch_to_scheduler(tokenized_obj)
|
||||
self._mark_state_dispatched(tokenized_obj.rid)
|
||||
dispatched = True
|
||||
tokenized_obj.time_stats = time_stats
|
||||
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
||||
@@ -1578,6 +1585,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
if not dispatched:
|
||||
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(
|
||||
self,
|
||||
tokenized_objs: List[
|
||||
@@ -1605,6 +1622,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
||||
|
||||
self._dispatch_to_scheduler(batch_req)
|
||||
for tokenized_obj in tokenized_objs:
|
||||
self._mark_state_dispatched(tokenized_obj.rid)
|
||||
dispatched = True
|
||||
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
|
||||
tokenized_obj.time_stats = time_stat
|
||||
@@ -1819,7 +1838,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self,
|
||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||
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
|
||||
|
||||
generators = []
|
||||
@@ -1885,6 +1907,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
tokenized_obj.sampling_params.max_new_tokens = 0
|
||||
tokenized_obj.stream = False
|
||||
self._init_req_state(tmp_obj)
|
||||
request_rids.add(tmp_obj.rid)
|
||||
await self._send_one_request(tokenized_obj)
|
||||
await self._wait_one_response(tmp_obj, request).__anext__()
|
||||
|
||||
@@ -1901,6 +1924,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
]
|
||||
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
||||
self._init_req_state(tmp_obj)
|
||||
request_rids.add(tmp_obj.rid)
|
||||
state = self.rid_to_state[tmp_obj.rid]
|
||||
tokenized_obj.time_stats = state.time_stats
|
||||
if tmp_obj.return_prompt_token_ids:
|
||||
@@ -1970,14 +1994,21 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
if not abort_all and not rid:
|
||||
logger.warning("Ignore abort_request with empty rid and abort_all=False")
|
||||
return
|
||||
if (
|
||||
not abort_all
|
||||
and get_serving().tokenizer_worker_num == 1
|
||||
and rid not in self.rid_to_state
|
||||
):
|
||||
return
|
||||
state = None if abort_all else self.rid_to_state.get(rid)
|
||||
if not abort_all:
|
||||
if state is not None:
|
||||
if state.abort_sent:
|
||||
return
|
||||
state.abort_sent = True
|
||||
elif get_serving().tokenizer_worker_num == 1:
|
||||
return
|
||||
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:
|
||||
# TODO: also use custom_labels from the request
|
||||
self.metrics_collector.observe_one_aborted_request(
|
||||
@@ -2140,10 +2171,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Abort the request if the client is disconnected.
|
||||
async def abort_request():
|
||||
await asyncio.sleep(2)
|
||||
if obj.is_single:
|
||||
self.abort_request(obj.rid)
|
||||
else:
|
||||
for rid in obj.rid:
|
||||
rids = [obj.rid] if obj.is_single else obj.rid
|
||||
for rid in rids:
|
||||
if rid in self.rid_to_state:
|
||||
self.abort_request(rid)
|
||||
|
||||
background_tasks = BackgroundTasks()
|
||||
@@ -3455,19 +3485,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header)
|
||||
time_stats.set_created_time(created_time)
|
||||
|
||||
def _discard_pending_req_states(self, obj):
|
||||
"""Drop rid_to_state entries created by _init_req_state for *obj*.
|
||||
def _release_req_states_on_failure(self, rids: Iterable[str]):
|
||||
"""Release rid_to_state entries created for a failed handler.
|
||||
|
||||
Safe to call after a partial/failed dispatch: only entries still present
|
||||
are removed, and the scheduler-response path looks up state with
|
||||
``.get(...)`` so a later output for a discarded rid is ignored, not fatal.
|
||||
Undelivered states are removed locally. Dispatched requests are aborted
|
||||
and retained until the scheduler response removes them.
|
||||
"""
|
||||
if not hasattr(obj, "is_single") or obj.is_single:
|
||||
rids = [obj.rid]
|
||||
else:
|
||||
rids = obj.rid
|
||||
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(
|
||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||
|
||||
Reference in New Issue
Block a user