diff --git a/python/sglang/srt/disaggregation/fake/conn.py b/python/sglang/srt/disaggregation/fake/conn.py index 9fc18af80..10c0127a3 100644 --- a/python/sglang/srt/disaggregation/fake/conn.py +++ b/python/sglang/srt/disaggregation/fake/conn.py @@ -1,4 +1,5 @@ import logging +import time from typing import List, Optional import numpy as np @@ -13,6 +14,7 @@ from sglang.srt.disaggregation.base.conn import ( KVTransferMetric, ) from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.environ import envs from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -46,13 +48,26 @@ class FakeKVSender(BaseKVSender): req_has_disagg_prefill_dp_rank: bool = False, ): self.kv_mgr = mgr + self.bootstrap_room = bootstrap_room + # Set by any chunk, not only the last one: nothing is transferred, so a + # chunk that never comes cannot change the outcome. self.has_sent = False self.conclude_state: Optional[KVPoll] = None + # Read here rather than off kv_mgr: a FAKE_BOOTSTRAP_HOST req on a real + # backend pairs this sender with that backend's KVManager, which carries + # no waiting_timeout in prefill mode. + self.waiting_timeout = envs.SGLANG_DISAGGREGATION_WAITING_TIMEOUT.get() + self.inited = False + self.waiting_since: Optional[float] = None def poll(self) -> KVPoll: if self.conclude_state is not None: return self.conclude_state + if not self.has_sent: + timeout_result = self._check_waiting_timeout() + if timeout_result is not None: + return timeout_result # Assume handshake completed instantly return KVPoll.WaitingForInput @@ -61,18 +76,47 @@ class FakeKVSender(BaseKVSender): self.conclude_state = KVPoll.Success return KVPoll.Success + def _check_waiting_timeout(self) -> Optional[KVPoll]: + # A send() that never comes must not pin the prefill inflight queue forever. + # No deadline before init(): the request is still in the bootstrap queue. + if not self.inited: + return None + if self.waiting_since is None: + # Clock starts at the first poll after init(), not at init() itself: + # the scheduler stops polling while a request queues and computes + # prefill. Monotonic, so an NTP step cannot fail a healthy request. + self.waiting_since = time.monotonic() + return None + elapsed = time.monotonic() - self.waiting_since + if elapsed < self.waiting_timeout: + return None + logger.warning_once( + "Some FakeKVSender requests fail to receive a KV chunk after bootstrapping. " + "If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_WAITING_TIMEOUT=600' (10 minutes) to relax the timeout condition. " + ) + logger.debug( + f"FakeKVSender for room {self.bootstrap_room} timed out after {elapsed:.1f}s " + "in KVPoll.WaitingForInput; no KV chunk was ever sent." + ) + self.conclude_state = KVPoll.Failed + return KVPoll.Failed + def get_transfer_metric(self) -> KVTransferMetric: return KVTransferMetric() def init( self, - kv_indices: list[int], + num_kv_indices: int, aux_index: Optional[int] = None, ): + self.inited = True logger.debug( - f"FakeKVSender init with kv_indices: {kv_indices}, aux_index: {aux_index}" + f"FakeKVSender init with num_kv_indices: {num_kv_indices}, aux_index: {aux_index}" ) - pass + + def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool: + # A zero-page last chunk must still send: poll() only concludes after send(). + return num_pages > 0 or last_chunk def send( self, @@ -112,6 +156,9 @@ class FakeKVReceiver(BaseKVReceiver): if not self.bootstrap_done: return KVPoll.Bootstrapping if not self.has_sent_metadata: + # No deadline needed here, unlike FakeKVSender: send_metadata() is + # unconditional once the decode side preallocates, and waiting for + # KV space is not a stalled transfer. return KVPoll.WaitingForInput logger.debug("FakeKVReceiver poll success") self.conclude_state = KVPoll.Success diff --git a/test/registered/unit/disaggregation/test_fake_kv_sender.py b/test/registered/unit/disaggregation/test_fake_kv_sender.py new file mode 100644 index 000000000..f8949bb91 --- /dev/null +++ b/test/registered/unit/disaggregation/test_fake_kv_sender.py @@ -0,0 +1,162 @@ +import time +import unittest +from unittest.mock import MagicMock + +import numpy as np + +from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll +from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVSender +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.environ import envs +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +class TestFakeKVSender(unittest.TestCase): + """A FakeKVSender whose send() never comes must not pin the prefill + inflight queue forever, and a healthy request must never be failed early.""" + + def setUp(self): + self.mgr = self._make_mgr() + self.sender = self._make_sender() + + def _make_mgr(self) -> FakeKVManager: + return FakeKVManager( + args=KVArgs(), + disaggregation_mode=DisaggregationMode.PREFILL, + server_args=MagicMock(), + ) + + def _make_sender(self, mgr=None) -> FakeKVSender: + return FakeKVSender( + mgr=self.mgr if mgr is None else mgr, + bootstrap_addr="fake_addr:1234", + bootstrap_room=42, + dest_tp_ranks=[0], + pp_rank=0, + ) + + def _expire_deadline(self, sender: FakeKVSender, by: float = 1.0) -> None: + """Push the armed deadline `by` seconds past the timeout.""" + sender.waiting_since -= sender.waiting_timeout + by + + def test_polls_before_send_stay_waiting_for_input(self): + """Repeated polling must not conclude on its own: the scheduler polls a + waiting request once per loop iteration, and those are healthy requests.""" + self.sender.init(3, 5) + for _ in range(1000): + self.assertEqual(self.sender.poll(), KVPoll.WaitingForInput) + self.assertIsNone(self.sender.conclude_state) + + def test_normal_send_flow(self): + self.sender.init(3, 5) + self.assertEqual(self.sender.poll(), KVPoll.WaitingForInput) + + self.sender.send(np.array([0, 1, 2], dtype=np.int32)) + self.assertTrue(self.sender.has_sent) + + self.assertEqual(self.sender.poll(), KVPoll.Success) + self.assertEqual(self.sender.conclude_state, KVPoll.Success) + # Cached afterwards. + self.assertEqual(self.sender.poll(), KVPoll.Success) + + def test_zero_page_last_chunk_still_sends(self): + """A zero-page last chunk must not be gated out: skipping it leaves the + sender in WaitingForInput forever and pins the prefill inflight queue.""" + self.assertTrue(self.sender.should_send_kv_chunk(0, last_chunk=True)) + self.assertFalse(self.sender.should_send_kv_chunk(0, last_chunk=False)) + self.assertTrue(self.sender.should_send_kv_chunk(3, last_chunk=False)) + + def test_fully_cached_request_concludes(self): + """Regression: a request whose last chunk carries no pages still reaches + a terminal poll state instead of accumulating in the inflight queue.""" + self.sender.init(0, 0) + page_indices = np.array([], dtype=np.int32) + # Mirrors the send_kv_chunk gate in SchedulerDisaggregationPrefillMixin. + if self.sender.should_send_kv_chunk(len(page_indices), True): + self.sender.send(page_indices) + self.assertEqual(self.sender.poll(), KVPoll.Success) + + def test_send_without_init_or_poll(self): + self.sender.send(np.array([0, 1], dtype=np.int32)) + self.assertEqual(self.sender.poll(), KVPoll.Success) + + def test_waiting_timeout_fails_a_stuck_sender(self): + self.sender.init(1, 0) + self.assertEqual(self.sender.poll(), KVPoll.WaitingForInput) + + self._expire_deadline(self.sender) + self.assertEqual(self.sender.poll(), KVPoll.Failed) + self.assertEqual(self.sender.conclude_state, KVPoll.Failed) + # Terminal: stays Failed even if send() arrives late. + self.sender.send(np.array([0], dtype=np.int32)) + self.assertEqual(self.sender.poll(), KVPoll.Failed) + + def test_poll_does_not_read_the_timeout_off_the_manager(self): + """Regression: a FAKE_BOOTSTRAP_HOST req on a real transfer backend pairs + FakeKVSender with that backend's KVManager, which defines no + waiting_timeout in prefill mode. Reading the knob off the manager raised + AttributeError in the scheduler loop on the first poll after init().""" + + class ManagerWithoutWaitingTimeout: + pass + + sender = self._make_sender(mgr=ManagerWithoutWaitingTimeout()) + sender.init(1, 0) + self.assertEqual(sender.poll(), KVPoll.WaitingForInput) + + self._expire_deadline(sender) + self.assertEqual(sender.poll(), KVPoll.Failed) + + def test_queue_and_prefill_time_is_not_charged_to_the_deadline(self): + """The scheduler stops polling between init() and the last chunk, so the + deadline starts at the first poll. Arming it at init() instead would fail + a healthy request that merely queued for longer than the timeout.""" + self.sender.init(1, 0) + self.assertIsNone(self.sender.waiting_since) + + after_init = time.monotonic() + self.assertEqual(self.sender.poll(), KVPoll.WaitingForInput) + self.assertGreaterEqual(self.sender.waiting_since, after_init) + + def test_no_timeout_before_init(self): + """A sender still in the bootstrap queue has no deadline, matching the + real backends, which cover that window with the bootstrap timeout.""" + self.assertFalse(self.sender.inited) + for _ in range(100): + self.assertEqual(self.sender.poll(), KVPoll.WaitingForInput) + self.assertIsNone(self.sender.waiting_since) + + def test_timeout_is_configurable(self): + """The knob is read once, when the sender is built.""" + with envs.SGLANG_DISAGGREGATION_WAITING_TIMEOUT.override(600): + sender = self._make_sender() + self.assertEqual(sender.waiting_timeout, 600) + + sender.init(1, 0) + self.assertEqual(sender.poll(), KVPoll.WaitingForInput) + sender.waiting_since -= 599 + self.assertEqual(sender.poll(), KVPoll.WaitingForInput) + sender.waiting_since -= 2 + self.assertEqual(sender.poll(), KVPoll.Failed) + + def test_abort_sets_failed_state(self): + self.sender.abort() + self.assertEqual(self.sender.conclude_state, KVPoll.Failed) + self.assertEqual(self.sender.poll(), KVPoll.Failed) + + def test_get_transfer_metric(self): + metric = self.sender.get_transfer_metric() + self.assertIsNone(metric.transfer_latency_s) + self.assertIsNone(metric.alloc_latency_s) + self.assertIsNone(metric.transfer_total_bytes) + + def test_failure_exception(self): + with self.assertRaises(Exception) as ctx: + self.sender.failure_exception() + self.assertIn("Fake KVSender Exception", str(ctx.exception)) + + +if __name__ == "__main__": + unittest.main()