"""Unit tests for the KV-transfer metric accounting (port of upstream #24416). These exercise the pure metric logic without standing up a real PD transfer: - bytes are counted from the *actually-sent* indices (not prompt length), - duration_between guards uninitialized (<=0) timestamps, - compute_and_observe_kv_transfer_metrics consumes a KVTransferMetric. """ import unittest import numpy as np from sglang.srt.disaggregation.base.conn import KVTransferMetric from sglang.srt.disaggregation.common.conn import CommonKVSender from sglang.srt.observability.req_time_stats import SchedulerReqTimeStats class _FakeMgr: def __init__(self, kv_sum, state_sum): self.kv_item_lens_sum = kv_sum self.state_item_lens_sum = state_sum def _make_sender(kv_sum, state_sum): s = CommonKVSender.__new__(CommonKVSender) # bypass heavy __init__ s.kv_mgr = _FakeMgr(kv_sum, state_sum) s._transfer_metric = KVTransferMetric() s._transfer_num_kv_indices = 0 s._transfer_num_state_indices = 0 return s def _make_stats(): st = SchedulerReqTimeStats.__new__(SchedulerReqTimeStats) st.enable_metrics = False st.transfer_total_mb = 0.0 st.transfer_speed_gb_s = 0.0 st.prefill_transfer_queue_entry_time = 0.0 st.completion_time = 0.0 st.prefill_bootstrap_queue_entry_time = 0.0 st.bootstrap_done_time = 0.0 st.wait_queue_entry_time = 0.0 return st class TestKVTransferMetricBytes(unittest.TestCase): def test_records_actually_sent_indices(self): s = _make_sender(kv_sum=100, state_sum=10) s._record_transfer_indices(np.arange(5, dtype=np.int32), [1, 2, 3]) s._record_transfer_indices(np.arange(2, dtype=np.int32), None) m = s.get_transfer_metric() # 7 kv pages * 100 + 3 state pages * 10 self.assertEqual(m.transfer_total_bytes, 7 * 100 + 3 * 10) def test_no_state_pool(self): s = _make_sender(kv_sum=64, state_sum=0) s._record_transfer_indices(np.arange(4, dtype=np.int32), None) self.assertEqual(s.get_transfer_metric().transfer_total_bytes, 4 * 64) def test_bytes_track_sent_pages_not_prompt_length(self): # The whole point of #24416: under cache hits / CP shared-KV filtering only # a subset of pages is sent, and the metric must reflect that subset. s = _make_sender(kv_sum=100, state_sum=0) s._record_transfer_indices(np.arange(3, dtype=np.int32), None) self.assertEqual(s.get_transfer_metric().transfer_total_bytes, 3 * 100) class TestDurationBetween(unittest.TestCase): def test_guards_nonpositive(self): st = _make_stats() self.assertEqual(st.duration_between(0, 5), 0.0) self.assertEqual(st.duration_between(5, 0), 0.0) self.assertEqual(st.duration_between(-1.0, 5.0), 0.0) def test_positive(self): st = _make_stats() self.assertAlmostEqual(st.duration_between(2.0, 5.0), 3.0) class TestComputeAndObserve(unittest.TestCase): def test_none_bytes_returns_none(self): st = _make_stats() m = KVTransferMetric(transfer_latency_s=1.0, transfer_total_bytes=None) self.assertIsNone(st.compute_and_observe_kv_transfer_metrics(m)) def test_explicit_latency_and_bytes(self): st = _make_stats() mb = 1024 * 1024 m = KVTransferMetric(transfer_latency_s=2.0, transfer_total_bytes=20 * mb) out = st.compute_and_observe_kv_transfer_metrics(m) self.assertAlmostEqual(out["total_mb"], 20.0) self.assertAlmostEqual(out["latency_ms"], 2000.0) self.assertAlmostEqual(out["speed_gb_s"], (20.0 / 1024) / 2.0) def test_latency_derived_from_timestamps(self): st = _make_stats() st.prefill_transfer_queue_entry_time = 1.0 st.completion_time = 1.5 m = KVTransferMetric(transfer_latency_s=None, transfer_total_bytes=1024 * 1024) out = st.compute_and_observe_kv_transfer_metrics(m) self.assertAlmostEqual(out["latency_ms"], 500.0) def test_missing_timestamps_returns_none(self): st = _make_stats() # both timestamps are 0 m = KVTransferMetric(transfer_latency_s=None, transfer_total_bytes=1024) self.assertIsNone(st.compute_and_observe_kv_transfer_metrics(m)) class TestSendersImplementMetric(unittest.TestCase): def test_senders_have_get_transfer_metric(self): from sglang.srt.disaggregation.fake.conn import FakeKVSender from sglang.srt.disaggregation.mooncake.conn import MooncakeKVSender for cls in (CommonKVSender, FakeKVSender, MooncakeKVSender): self.assertTrue( callable(getattr(cls, "get_transfer_metric", None)), f"{cls.__name__} missing get_transfer_metric", ) if __name__ == "__main__": unittest.main()