From 7976928c5776b8519290c8cc9b1e313b7aeb3a95 Mon Sep 17 00:00:00 2001 From: cctry Date: Wed, 17 Jun 2026 23:13:28 -0700 Subject: [PATCH] Abort during chunked prefill + PD peer-liveness abort (#28086) --- python/sglang/srt/disaggregation/prefill.py | 39 +++ python/sglang/srt/managers/scheduler.py | 53 +++ .../batch_result_processor.py | 8 +- .../test_chunked_prefill_abort.py | 134 ++++++++ ...st_disaggregation_chunked_prefill_abort.py | 109 ++++++ ...test_disaggregation_peer_liveness_abort.py | 316 ++++++++++++++++++ .../test_disaggregation_basic.py | 106 +++++- .../test_scheduler_chunked_req_gate.py | 1 + 8 files changed, 763 insertions(+), 3 deletions(-) create mode 100644 python/sglang/test/manual/disaggregation/test_chunked_prefill_abort.py create mode 100644 python/sglang/test/manual/disaggregation/test_disaggregation_chunked_prefill_abort.py create mode 100644 python/sglang/test/manual/disaggregation/test_disaggregation_peer_liveness_abort.py diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index a57bedb36..6f244ff10 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -406,16 +406,55 @@ class SchedulerDisaggregationPrefillMixin: if room is not None and room in kv_mgr.transfer_infos: prefetch(room) + def resolve_waiting_queue_bootstrap(self: Scheduler) -> None: + """Resolve bootstrap status for waiting prefill requests before admission. + + Covers the window between leaving the bootstrap queue and being admitted + into a running batch: aborts requests whose decode peer died, and + finalizes optimistic requests whose bootstrap completed so they skip + the post-forward bootstrap check. + """ + candidates = [req for req in self.waiting_queue if not is_aborted(req)] + if not candidates: + return + polls = poll_and_all_reduce_attn_cp_tp_group( + [req.disagg_kv_sender for req in candidates], + self.attn_cp_cpu_group, + self.attn_tp_cpu_group, + ) + failed = set() + for req, poll in zip(candidates, polls): + if poll == KVPoll.Failed: + self.handle_bootstrap_failure(req) + failed.add(req) + elif ( + poll == KVPoll.WaitingForInput + and req.pending_bootstrap + and not should_force_retry(req) + ): + # Optimistic requests reserved a metadata buffer when popped, so + # finalize cannot fail here; if it ever does, the request stays + # pending and the post-forward check resolves it. + self.disagg_prefill_bootstrap_queue.finalize_bootstrap(req) + if failed: + self.waiting_queue = [ + req for req in self.waiting_queue if req not in failed + ] + @scheduler_nvtx_method("scheduler.get_next_batch_to_run") def get_next_disagg_prefill_batch_to_run( self: Scheduler, ) -> Optional[ScheduleBatch]: + self.process_pending_chunked_abort() + # HACK (byronhsu): reset the batch_is_full flag because we never enter update_running_batch which resets it # Otherwise, it hangs under high concurrency self.running_batch.batch_is_full = False self.process_prefill_chunk() + self.resolve_waiting_queue_bootstrap() + batch = self.get_new_batch_prefill() batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f316aef42..1057626b5 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -985,6 +985,7 @@ class Scheduler( elif self.chunked_prefill_size is not None and self.chunked_prefill_size <= 0: self.chunked_prefill_size = None self.chunked_req = None + self._pending_chunked_abort_req = None self.is_mixed_chunk = ( self.chunked_prefill_size is not None and self.server_args.enable_mixed_chunk @@ -2424,6 +2425,52 @@ class Scheduler( def stash_chunked_request(self, req: Req): maybe_cache_unfinished_req(req, self.tree_cache, chunked=True) + def process_pending_chunked_abort(self) -> None: + """Abort an in-flight chunked-prefill request once it is safe to do so. + + ``abort_request`` only records the target in ``_pending_chunked_abort_req`` + (tearing it down mid-iteration is unsafe). Clearing ``chunked_req`` here at + the top of the scheduling step stops the next chunk from launching; the + chunk already launched is drained when its result is resolved. Under overlap + the result lands a step later, so the batch-result processors keep + ``inflight_middle_chunks`` accounting intact and skip the aborted chunk: + ``process_batch_result_disagg_prefill`` via its ``is_aborted`` drop, and + ``process_batch_result_prefill`` via its chunked branch (the finished req + is excluded from streaming and its logprob offset is still accounted). + Mirrors ``handle_bootstrap_failure``. + """ + req = self._pending_chunked_abort_req + if req is None: + return + if self.chunked_req is not req: + # Already past chunked prefill; the running-batch abort path handles + # it. Drop the marker once the request is actually gone. + if req.finished() or req.req_pool_idx is None: + self._pending_chunked_abort_req = None + return + + prepare_abort(req, "Aborted") + req.time_stats.trace_ctx.abort(abort_info={"reason": "Aborted"}) + req.to_finish = None + if self.disaggregation_mode == DisaggregationMode.PREFILL: + req.disagg_kv_sender.abort() + maybe_release_metadata_buffer( + req, self.req_to_metadata_buffer_idx_allocator + ) + req.pending_bootstrap = False + if self.enable_hicache_storage: + self.tree_cache.release_aborted_request(req.rid) + if ( + req.req_pool_idx is not None or self.tree_cache.supports_mamba() + ) and not req.kv_committed_freed: + release_kv_cache(req, self.tree_cache, is_insert=False) + + self.chunked_req = None + self._chunked_req_scheduled_last_iter = False + self._pending_chunked_abort_req = None + self.ipc_channels.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req) + logger.debug(f"Abort chunked prefill request. {req.rid=}") + def _build_hisparse_decode_batch(self, reqs): """Build a ScheduleBatch for hisparse requests transitioning from staging to decode.""" device = self.device @@ -2467,6 +2514,8 @@ class Scheduler( @scheduler_nvtx_method("scheduler.get_next_batch_to_run") def get_next_batch_to_run(self) -> Optional[ScheduleBatch]: + self.process_pending_chunked_abort() + if self.enable_fpm: self._fpm_batch_t0 = time.monotonic() self._abort_on_waiting_timeout() @@ -3684,6 +3733,10 @@ class Scheduler( return RpcReqOutput(success, "" if not exec else str(exec)) def abort_request(self, recv_req: AbortReq): + if (chunked_req := self.chunked_req) is not None: + if recv_req.abort_all or chunked_req.rid.startswith(recv_req.rid): + self._pending_chunked_abort_req = chunked_req + # todo hisparse, release resources for abort requests in hisparse coordinator # Delete requests in the waiting queue to_del = [] diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 7ac74073f..a636bf201 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -216,8 +216,12 @@ class SchedulerBatchResultProcessor: logprob_pt = 0 for i, (req, next_token_id) in enumerate(zip(batch.reqs, next_token_ids)): - if req.finished() or req.is_retracted: - # decode req in mixed batch or retracted req + if ( + req.finished() and req.inflight_middle_chunks <= 0 + ) or req.is_retracted: + # Decode req in a mixed batch, or a retracted req. Keep an + # aborted middle chunk in the chunked branch long enough to + # drain its accounting without streaming it. continue if req.inflight_middle_chunks <= 0: diff --git a/python/sglang/test/manual/disaggregation/test_chunked_prefill_abort.py b/python/sglang/test/manual/disaggregation/test_chunked_prefill_abort.py new file mode 100644 index 000000000..af7f47fe7 --- /dev/null +++ b/python/sglang/test/manual/disaggregation/test_chunked_prefill_abort.py @@ -0,0 +1,134 @@ +import threading +import time +import unittest +import uuid +from typing import Any + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +CHUNKED_PREFILL_SIZE = 64 +LONG_PROMPT = ( + "The quick brown fox jumps over the lazy dog. " + "Pack my box with five dozen liquor jugs. " + "Sphinx of black quartz, judge my vow. " +) * 900 + + +def _decode_response(response: requests.Response) -> Any: + try: + return response.json() + except ValueError: + return response.text + + +def _is_abort_result(status_code: int, body: Any) -> bool: + if status_code == 200: + reason = ( + body.get("meta_info", {}).get("finish_reason", {}) + if isinstance(body, dict) + else {} + ) + return isinstance(reason, dict) and reason.get("type") == "abort" + + if status_code not in (500, 503): + return False + + text = body if isinstance(body, str) else str(body) + return "abort" in text.lower() + + +class TestChunkedPrefillAbortE2E(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--disable-cuda-graph", + "--chunked-prefill-size", + str(CHUNKED_PREFILL_SIZE), + "--max-running-requests", + "4", + ], + ) + + @classmethod + def tearDownClass(cls): + if cls.process: + kill_process_tree(cls.process.pid) + + def test_abort_mid_chunked_prefill_by_rid(self): + rid = f"chunked-prefill-abort-{uuid.uuid4().hex}" + result: dict[str, Any] = {} + + def run_generate(): + try: + response = requests.post( + self.base_url + "/generate", + json={ + "rid": rid, + "text": f"{rid}\n{LONG_PROMPT}", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 4096, + "ignore_eos": True, + }, + }, + timeout=180, + ) + result["status_code"] = response.status_code + result["body"] = _decode_response(response) + except requests.RequestException as exc: + result["exception"] = repr(exc) + + thread = threading.Thread(target=run_generate) + thread.start() + + time.sleep(0.5) + abort_deadline = time.monotonic() + 8 + while thread.is_alive() and time.monotonic() < abort_deadline: + requests.post( + self.base_url + "/abort_request", + json={"rid": rid, "abort_all": False}, + timeout=10, + ) + time.sleep(0.2) + + thread.join(timeout=60) + self.assertFalse(thread.is_alive(), "Chunked-prefill abort request hung") + self.assertNotIn("exception", result, result.get("exception")) + self.assertTrue( + _is_abort_result(result["status_code"], result["body"]), + f"Expected chunked-prefill request to abort, got {result}", + ) + + health = requests.get(self.base_url + "/health", timeout=10) + self.assertEqual(health.status_code, 200, health.text) + + follow_up = requests.post( + self.base_url + "/generate", + json={ + "rid": f"chunked-prefill-after-abort-{uuid.uuid4().hex}", + "text": "The capital of France is", + "sampling_params": {"temperature": 0, "max_new_tokens": 4}, + }, + timeout=60, + ) + follow_up_body = _decode_response(follow_up) + self.assertEqual(follow_up.status_code, 200, follow_up_body) + self.assertFalse(_is_abort_result(follow_up.status_code, follow_up_body)) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/test/manual/disaggregation/test_disaggregation_chunked_prefill_abort.py b/python/sglang/test/manual/disaggregation/test_disaggregation_chunked_prefill_abort.py new file mode 100644 index 000000000..44a3643e7 --- /dev/null +++ b/python/sglang/test/manual/disaggregation/test_disaggregation_chunked_prefill_abort.py @@ -0,0 +1,109 @@ +import threading +import time +import unittest +import uuid +from typing import Any + +import requests + +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) +from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST + +PD_EXTRA_ARGS = ["--max-running-requests", "4", "--chunked-prefill-size", "64"] +LONG_PROMPT = ( + "The quick brown fox jumps over the lazy dog. " + "Pack my box with five dozen liquor jugs. " + "Sphinx of black quartz, judge my vow. " +) * 900 + + +def _decode_response(response: requests.Response) -> Any: + try: + return response.json() + except ValueError: + return response.text + + +def _is_abort_result(status_code: int, body: Any) -> bool: + if status_code == 200: + reason = ( + body.get("meta_info", {}).get("finish_reason", {}) + if isinstance(body, dict) + else {} + ) + return isinstance(reason, dict) and reason.get("type") == "abort" + + if status_code not in (500, 503): + return False + + text = body if isinstance(body, str) else str(body) + return "abort" in text.lower() + + +class TestDisaggChunkedPrefillAbort(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST + cls.extra_prefill_args = PD_EXTRA_ARGS + cls.extra_decode_args = PD_EXTRA_ARGS + cls.launch_all() + + def _post_abort(self, rid: str): + for url in (self.prefill_url, self.decode_url): + requests.post( + url + "/abort_request", + json={"rid": rid, "abort_all": False}, + timeout=10, + ) + + def test_abort_mid_chunked_prefill_by_rid(self): + rid = f"pd-chunked-prefill-abort-{uuid.uuid4().hex}" + result: dict[str, Any] = {} + + def run_generate(): + try: + response = requests.post( + self.lb_url + "/generate", + json={ + "rid": rid, + "text": f"{rid}\n{LONG_PROMPT}", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 4096, + "ignore_eos": True, + }, + }, + timeout=180, + ) + result["status_code"] = response.status_code + result["body"] = _decode_response(response) + except requests.RequestException as exc: + result["exception"] = repr(exc) + + thread = threading.Thread(target=run_generate) + thread.start() + + time.sleep(1.0) + abort_deadline = time.monotonic() + 8 + while thread.is_alive() and time.monotonic() < abort_deadline: + self._post_abort(rid) + time.sleep(0.2) + + thread.join(timeout=60) + self.assertFalse(thread.is_alive(), "Chunked-prefill abort request hung") + self.assertNotIn("exception", result, result.get("exception")) + self.assertTrue( + _is_abort_result(result["status_code"], result["body"]), + f"Expected chunked-prefill request to abort, got {result}", + ) + + for url in (self.lb_url, self.prefill_url, self.decode_url): + health = requests.get(url + "/health", timeout=10) + self.assertEqual(health.status_code, 200, health.text) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/test/manual/disaggregation/test_disaggregation_peer_liveness_abort.py b/python/sglang/test/manual/disaggregation/test_disaggregation_peer_liveness_abort.py new file mode 100644 index 000000000..63696c903 --- /dev/null +++ b/python/sglang/test/manual/disaggregation/test_disaggregation_peer_liveness_abort.py @@ -0,0 +1,316 @@ +import json +import threading +import time +import unittest +import uuid +from concurrent.futures import ThreadPoolExecutor, TimeoutError, as_completed +from typing import Any + +import requests + +from sglang.srt.environ import envs +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) +from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST + +FAILURE_PROB = 0.12 +PD_EXTRA_ARGS = [ + "--max-running-requests", + "4", + "--chunked-prefill-size", + "128", +] +PD_LONG_PROMPT = ( + "The quick brown fox jumps over the lazy dog. " + "Pack my box with five dozen liquor jugs. " + "Sphinx of black quartz, judge my vow. " +) * 240 + + +def _body_text(body: Any) -> str: + if isinstance(body, str): + return body + return json.dumps(body, sort_keys=True) + + +def _decode_response(response: requests.Response) -> Any: + try: + return response.json() + except ValueError: + return response.text + + +def _finish_reason(body: Any) -> dict[str, Any]: + if not isinstance(body, dict): + return {} + reason = body.get("meta_info", {}).get("finish_reason", {}) + return reason if isinstance(reason, dict) else {} + + +def _is_abort_result(status_code: int, body: Any) -> bool: + if status_code == 200: + return _finish_reason(body).get("type") == "abort" + + if status_code not in (500, 503): + return False + + text = _body_text(body).lower() + return any( + marker in text + for marker in ( + "pd peer failed", + "prefill bootstrap failed", + "abort", + "aborted", + ) + ) + + +def _is_success_result(status_code: int, body: Any) -> bool: + return ( + status_code == 200 + and isinstance(body, dict) + and "text" in body + and _finish_reason(body).get("type") != "abort" + ) + + +class TestDisaggregationPeerLivenessAbort(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls._disagg_failure_ctx = envs.SGLANG_TEST_DISAGG_FAILURE_PROB.override( + FAILURE_PROB + ) + cls._disagg_failure_ctx.__enter__() + + cls.model = DEFAULT_MODEL_NAME_FOR_TEST + cls.extra_prefill_args = PD_EXTRA_ARGS + cls.extra_decode_args = PD_EXTRA_ARGS + cls.launch_all() + + @classmethod + def tearDownClass(cls): + try: + super().tearDownClass() + finally: + if cls._disagg_failure_ctx: + cls._disagg_failure_ctx.__exit__(None, None, None) + + def _post_generate(self, rid: str, max_new_tokens: int = 32) -> dict[str, Any]: + response = requests.post( + self.lb_url + "/generate", + json={ + "rid": rid, + "text": f"{rid}\n{PD_LONG_PROMPT}", + "sampling_params": { + "temperature": 0, + "max_new_tokens": max_new_tokens, + "ignore_eos": True, + }, + }, + timeout=75, + ) + return { + "rid": rid, + "status_code": response.status_code, + "body": _decode_response(response), + } + + def _assert_servers_healthy(self): + for url in (self.lb_url, self.prefill_url, self.decode_url): + response = requests.get(url + "/health", timeout=10) + self.assertEqual(response.status_code, 200, response.text) + + def _load_count(self, url: str) -> int: + response = requests.get( + url + "/v1/loads?include=core,disagg,queues", + timeout=10, + ) + response.raise_for_status() + total = 0 + for load in response.json()["loads"]: + total += int(load.get("num_running_reqs", 0)) + total += int(load.get("num_waiting_reqs", 0)) + disagg = load.get("disaggregation", {}) + total += int(disagg.get("prefill_bootstrap_queue_reqs", 0)) + total += int(disagg.get("prefill_inflight_queue_reqs", 0)) + total += int(disagg.get("decode_prealloc_queue_reqs", 0)) + total += int(disagg.get("decode_transfer_queue_reqs", 0)) + total += int(disagg.get("decode_retracted_queue_reqs", 0)) + queues = load.get("queues", {}) + total += int(queues.get("waiting", 0)) + total += int(queues.get("grammar", 0)) + total += int(queues.get("paused", 0)) + total += int(queues.get("retracted", 0)) + return total + + def _assert_eventually_idle(self): + deadline = time.monotonic() + 30 + last_counts = {} + while time.monotonic() < deadline: + last_counts = { + self.prefill_url: self._load_count(self.prefill_url), + self.decode_url: self._load_count(self.decode_url), + } + if all(count == 0 for count in last_counts.values()): + return + time.sleep(1) + + self.fail(f"PD servers still have request load after cleanup: {last_counts}") + + def _assert_one_successful_generate_with_retries(self): + last_result = None + for _ in range(20): + rid = f"pd-peer-liveness-health-{uuid.uuid4().hex}" + last_result = self._post_generate(rid, max_new_tokens=1) + if _is_success_result(last_result["status_code"], last_result["body"]): + return + time.sleep(0.2) + + self.fail( + "Server stayed HTTP-healthy but did not complete a follow-up generate " + f"under failure injection. Last result: {last_result}" + ) + + def test_peer_failures_abort_without_hanging(self): + num_requests = 36 + start = time.monotonic() + results: list[dict[str, Any]] = [] + + with ThreadPoolExecutor(max_workers=18) as executor: + futures = [ + executor.submit( + self._post_generate, + f"pd-peer-liveness-{i}-{uuid.uuid4().hex}", + ) + for i in range(num_requests) + ] + + try: + for future in as_completed(futures, timeout=120): + results.append(future.result()) + except TimeoutError: + unfinished = sum(not future.done() for future in futures) + self.fail(f"{unfinished} generate requests did not return in time") + + elapsed = time.monotonic() - start + self.assertLess(elapsed, 120) + self.assertEqual(len(results), num_requests) + + aborts = [ + result + for result in results + if _is_abort_result(result["status_code"], result["body"]) + ] + successes = [ + result + for result in results + if _is_success_result(result["status_code"], result["body"]) + ] + unexpected = [ + result + for result in results + if result not in aborts and result not in successes + ] + + self.assertFalse( + unexpected, + "Expected only clean aborts or successful completions. " + f"Unexpected results: {unexpected[:3]}", + ) + self.assertGreater( + len(aborts), + 0, + "Injected KVPoll.Failed should abort at least one request", + ) + + self._assert_servers_healthy() + self._assert_eventually_idle() + self._assert_one_successful_generate_with_retries() + + +class TestDisaggregationPeerLivenessCleanRecovery(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = DEFAULT_MODEL_NAME_FOR_TEST + cls.extra_prefill_args = PD_EXTRA_ARGS + cls.extra_decode_args = PD_EXTRA_ARGS + cls.launch_all() + + def _post_abort_to_engines(self, rid: str): + for url in (self.prefill_url, self.decode_url): + response = requests.post( + url + "/abort_request", + json={"rid": rid, "abort_all": False}, + timeout=10, + ) + self.assertEqual(response.status_code, 200, response.text) + + def _assert_servers_healthy(self): + for url in (self.lb_url, self.prefill_url, self.decode_url): + response = requests.get(url + "/health", timeout=10) + self.assertEqual(response.status_code, 200, response.text) + + def test_clean_generate_and_explicit_abort_by_rid(self): + clean_response = requests.post( + self.lb_url + "/generate", + json={ + "rid": f"pd-clean-generate-{uuid.uuid4().hex}", + "text": "The capital of France is", + "sampling_params": {"temperature": 0, "max_new_tokens": 4}, + }, + timeout=60, + ) + clean_body = _decode_response(clean_response) + self.assertTrue( + _is_success_result(clean_response.status_code, clean_body), clean_body + ) + + rid = f"pd-explicit-abort-{uuid.uuid4().hex}" + result: dict[str, Any] = {} + + def run_generate(): + try: + response = requests.post( + self.lb_url + "/generate", + json={ + "rid": rid, + "text": f"{rid}\n{PD_LONG_PROMPT}", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 4096, + "ignore_eos": True, + }, + }, + timeout=180, + ) + result["status_code"] = response.status_code + result["body"] = _decode_response(response) + except requests.RequestException as exc: + result["exception"] = repr(exc) + + thread = threading.Thread(target=run_generate) + thread.start() + + time.sleep(0.25) + abort_deadline = time.monotonic() + 8 + while thread.is_alive() and time.monotonic() < abort_deadline: + self._post_abort_to_engines(rid) + time.sleep(0.25) + + thread.join(timeout=60) + self.assertFalse(thread.is_alive(), "Explicitly aborted PD request hung") + self.assertNotIn("exception", result, result.get("exception")) + self.assertTrue( + _is_abort_result(result["status_code"], result["body"]), + f"Expected explicit abort result, got {result}", + ) + + self._assert_servers_healthy() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/disaggregation/test_disaggregation_basic.py b/test/registered/disaggregation/test_disaggregation_basic.py index 048b3e014..74509a6d0 100644 --- a/test/registered/disaggregation/test_disaggregation_basic.py +++ b/test/registered/disaggregation/test_disaggregation_basic.py @@ -1,8 +1,12 @@ import asyncio import json import os +import threading +import time import unittest +import uuid from types import SimpleNamespace +from typing import Any import aiohttp import openai @@ -19,10 +23,11 @@ from sglang.test.server_fixtures.disaggregation_fixture import ( from sglang.test.test_utils import ( DEFAULT_DRAFT_MODEL_EAGLE3, DEFAULT_MODEL_NAME_FOR_TEST, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TARGET_MODEL_EAGLE3, ) -register_cuda_ci(est_time=560, stage="base-b", runner_config="2-gpu-large") +register_cuda_ci(est_time=700, stage="base-b", runner_config="2-gpu-large") class TestDisaggregationAccuracy(PauseResumeInPlaceMixin, PDDisaggregationServerBase): @@ -441,5 +446,104 @@ class TestDisaggregationPauseResumePrefillLeak(PDDisaggregationServerBase): ) +PD_CHUNKED_ABORT_EXTRA_ARGS = [ + "--max-running-requests", + "4", + "--chunked-prefill-size", + "64", +] +_CHUNKED_ABORT_LONG_PROMPT = ( + "The quick brown fox jumps over the lazy dog. " + "Pack my box with five dozen liquor jugs. " + "Sphinx of black quartz, judge my vow. " +) * 900 + + +def _decode_response(response: requests.Response) -> Any: + try: + return response.json() + except ValueError: + return response.text + + +def _is_abort_result(status_code: int, body: Any) -> bool: + if status_code == 200: + reason = ( + body.get("meta_info", {}).get("finish_reason", {}) + if isinstance(body, dict) + else {} + ) + return isinstance(reason, dict) and reason.get("type") == "abort" + + if status_code not in (500, 503): + return False + + text = body if isinstance(body, str) else str(body) + return "abort" in text.lower() + + +class TestDisaggChunkedPrefillAbort(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST + cls.extra_prefill_args = PD_CHUNKED_ABORT_EXTRA_ARGS + cls.extra_decode_args = PD_CHUNKED_ABORT_EXTRA_ARGS + cls.launch_all() + + def _post_abort(self, rid: str): + for url in (self.prefill_url, self.decode_url): + requests.post( + url + "/abort_request", + json={"rid": rid, "abort_all": False}, + timeout=10, + ) + + def test_abort_mid_chunked_prefill_by_rid(self): + rid = f"pd-chunked-prefill-abort-{uuid.uuid4().hex}" + result: dict[str, Any] = {} + + def run_generate(): + try: + response = requests.post( + self.lb_url + "/generate", + json={ + "rid": rid, + "text": f"{rid}\n{_CHUNKED_ABORT_LONG_PROMPT}", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 4096, + "ignore_eos": True, + }, + }, + timeout=180, + ) + result["status_code"] = response.status_code + result["body"] = _decode_response(response) + except requests.RequestException as exc: + result["exception"] = repr(exc) + + thread = threading.Thread(target=run_generate) + thread.start() + + time.sleep(1.0) + abort_deadline = time.monotonic() + 8 + while thread.is_alive() and time.monotonic() < abort_deadline: + self._post_abort(rid) + time.sleep(0.2) + + thread.join(timeout=60) + self.assertFalse(thread.is_alive(), "Chunked-prefill abort request hung") + self.assertNotIn("exception", result, result.get("exception")) + self.assertTrue( + _is_abort_result(result["status_code"], result["body"]), + f"Expected chunked-prefill request to abort, got {result}", + ) + + for url in (self.lb_url, self.prefill_url, self.decode_url): + health = requests.get(url + "/health", timeout=10) + self.assertEqual(health.status_code, 200, health.text) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/managers/test_scheduler_chunked_req_gate.py b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py index 4298907da..03f1efae0 100644 --- a/test/registered/unit/managers/test_scheduler_chunked_req_gate.py +++ b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py @@ -99,6 +99,7 @@ def _scheduler_for_get_next_batch(*, tree_cache, chunked_req) -> Scheduler: s.update_running_batch = MagicMock(side_effect=lambda batch: batch) s.tree_cache = tree_cache s.chunked_req = chunked_req + s._pending_chunked_abort_req = None return s