import unittest from collections import deque from types import SimpleNamespace from unittest.mock import MagicMock from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.managers.io_struct import HealthCheckOutput from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=2, suite="stage-a-test-cpu") class TestSchedulerHealthCheck(CustomTestCase): def _idle_scheduler(self, *, disaggregation_mode=DisaggregationMode.NULL): return SimpleNamespace( running_batch=SimpleNamespace(is_empty=lambda: True, reqs=[]), chunked_req=None, dllm_manager=SimpleNamespace(any_staging_reqs=lambda: False), last_batch=None, cur_batch=None, enable_overlap=False, result_queue=deque(), pp_size=1, waiting_queue=[], grammar_manager=SimpleNamespace(grammar_queue=[]), disaggregation_mode=disaggregation_mode, enable_hierarchical_cache=False, ) def test_decode_prealloc_queue_is_busy_for_health_check(self): scheduler = self._idle_scheduler(disaggregation_mode=DisaggregationMode.DECODE) scheduler.disagg_decode_prealloc_queue = SimpleNamespace(queue=[object()]) scheduler.disagg_decode_transfer_queue = SimpleNamespace(queue=[]) scheduler.server_args = SimpleNamespace( disaggregation_decode_enable_offload_kvcache=False ) self.assertFalse(Scheduler.is_fully_idle(scheduler, for_health_check=True)) def test_busy_health_check_sends_immediate_signal(self): sent = [] scheduler = SimpleNamespace( session_controller=SimpleNamespace(maybe_reap=MagicMock()), is_fully_idle=lambda for_health_check=False: False, send_to_tokenizer=SimpleNamespace( send_output=lambda output, recv_obj=None: sent.append((output, recv_obj)) ), return_health_check_ipcs=deque(), _request_dispatcher=MagicMock(), ) req = SimpleNamespace(rid="HEALTH_CHECK_unit", http_worker_ipc="ipc://unit") Scheduler.process_input_requests(scheduler, [req]) scheduler._request_dispatcher.assert_not_called() self.assertEqual(len(sent), 1) self.assertIsInstance(sent[0][0], HealthCheckOutput) self.assertEqual(sent[0][0].http_worker_ipc, "ipc://unit") self.assertIs(sent[0][1], req) def test_shutdown_pending_request_without_response_ts_aborts_cleanly(self): event = SimpleNamespace(set=MagicMock()) state = SimpleNamespace( finished=False, obj=SimpleNamespace(stream=True), text="", output_ids=[], out_list=[], event=event, ) tokenizer_manager = SimpleNamespace(rid_to_state={"rid": state}) TokenizerManager._finish_all_pending_requests_on_shutdown(tokenizer_manager) self.assertTrue(state.finished) self.assertEqual(len(state.out_list), 1) self.assertEqual( state.out_list[0]["meta_info"]["finish_reason"]["type"], "abort" ) event.set.assert_called_once() if __name__ == "__main__": unittest.main()