diff --git a/test/registered/unit/observability/test_cpu_monitor.py b/test/registered/unit/observability/test_cpu_monitor.py index e4d34d46d..3cb4e3bd5 100644 --- a/test/registered/unit/observability/test_cpu_monitor.py +++ b/test/registered/unit/observability/test_cpu_monitor.py @@ -1,5 +1,8 @@ +import threading import time import unittest +from collections import namedtuple +from unittest.mock import MagicMock, patch from sglang.test.ci.ci_register import register_cpu_ci @@ -34,5 +37,64 @@ class TestCpuMonitor(unittest.TestCase): self.assertGreater(value, 0) +class TestCpuMonitorMocked(unittest.TestCase): + """Fast, deterministic tests for start_cpu_monitor_thread using mocks.""" + + @patch("prometheus_client.Counter") + @patch("sglang.srt.observability.cpu_monitor.psutil.Process") + @patch("sglang.srt.observability.cpu_monitor.time.sleep") + def test_delta_calculation_over_two_iterations( + self, mock_sleep, MockProcess, MockCounter + ): + """Verify delta=(user_diff+system_diff) and last_times update across iterations.""" + from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread + + CpuTimes = namedtuple("CpuTimes", ["user", "system"]) + mock_process = MockProcess.return_value + mock_process.cpu_times.side_effect = [ + CpuTimes(user=1.0, system=0.5), # initial (L18) + CpuTimes(user=2.5, system=1.0), # iteration 1 (L22) + CpuTimes(user=4.0, system=2.0), # iteration 2 (L22) + ] + + # Allow 2 loop iterations, then stop the thread. + # Override threading.excepthook to suppress the pytest warning from + # the intentional exception used to terminate the monitor loop. + remaining = [2] + orig_hook = threading.excepthook + + def controlled_sleep(seconds): + if remaining[0] <= 0: + raise SystemExit + remaining[0] -= 1 + + mock_sleep.side_effect = controlled_sleep + threading.excepthook = lambda args: None + + mock_labeled = MagicMock() + MockCounter.return_value.labels.return_value = mock_labeled + + thread = start_cpu_monitor_thread("my_component", interval=3.0) + thread.join(timeout=1.0) + threading.excepthook = orig_hook + + # Thread is daemon (L29) + self.assertTrue(thread.daemon) + + # Sleep called with correct interval (L21) + mock_sleep.assert_called_with(3.0) + + # Counter labeled with component (L26) + MockCounter.return_value.labels.assert_called_with(component="my_component") + + # Delta calculation (L23-24) and counter increment (L26) + inc_calls = mock_labeled.inc.call_args_list + self.assertEqual(len(inc_calls), 2) + # Iteration 1: (2.5 - 1.0) + (1.0 - 0.5) = 2.0 + self.assertAlmostEqual(inc_calls[0].args[0], 2.0) + # Iteration 2: (4.0 - 2.5) + (2.0 - 1.0) = 2.5 (proves last_times updated) + self.assertAlmostEqual(inc_calls[1].args[0], 2.5) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/observability/test_func_timer.py b/test/registered/unit/observability/test_func_timer.py new file mode 100644 index 000000000..cb58c7ed0 --- /dev/null +++ b/test/registered/unit/observability/test_func_timer.py @@ -0,0 +1,82 @@ +"""Unit tests for func_timer.py — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import asyncio +import unittest +from unittest.mock import MagicMock, patch + +import sglang.srt.observability.func_timer as func_timer +from sglang.srt.observability.func_timer import enable_func_timer, time_func_latency + + +class TestFuncTimer(unittest.TestCase): + def setUp(self): + self.orig_enable = func_timer.enable_metrics + self.orig_latency = func_timer.FUNC_LATENCY + + def tearDown(self): + func_timer.enable_metrics = self.orig_enable + func_timer.FUNC_LATENCY = self.orig_latency + + @patch("prometheus_client.Histogram") + def test_enable_func_timer(self, MockHistogram): + """Sets enable_metrics and creates FUNC_LATENCY histogram.""" + enable_func_timer() + self.assertTrue(func_timer.enable_metrics) + self.assertIs(func_timer.FUNC_LATENCY, MockHistogram.return_value) + MockHistogram.assert_called_once() + + def test_sync_disabled(self): + """Sync function passes through when metrics disabled.""" + func_timer.enable_metrics = False + + @time_func_latency + def add(a, b): + return a + b + + self.assertEqual(add(2, 3), 5) + + def test_sync_enabled(self): + """Sync function timed with custom name when metrics enabled.""" + mock_histogram = MagicMock() + func_timer.enable_metrics = True + func_timer.FUNC_LATENCY = mock_histogram + + @time_func_latency(name="custom_op") + def add(a, b): + return a + b + + self.assertEqual(add(2, 3), 5) + mock_histogram.labels.assert_called_with(name="custom_op") + mock_histogram.labels().observe.assert_called_once() + + def test_async_disabled(self): + """Async function passes through when metrics disabled.""" + func_timer.enable_metrics = False + + @time_func_latency + async def add(a, b): + return a + b + + self.assertEqual(asyncio.run(add(2, 3)), 5) + + def test_async_enabled(self): + """Async function timed with default name when metrics enabled.""" + mock_histogram = MagicMock() + func_timer.enable_metrics = True + func_timer.FUNC_LATENCY = mock_histogram + + @time_func_latency + async def add(a, b): + return a + b + + self.assertEqual(asyncio.run(add(2, 3)), 5) + mock_histogram.labels.assert_called_with(name="add") + mock_histogram.labels().observe.assert_called_once() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/observability/test_label_transform.py b/test/registered/unit/observability/test_label_transform.py new file mode 100644 index 000000000..31edc9a85 --- /dev/null +++ b/test/registered/unit/observability/test_label_transform.py @@ -0,0 +1,39 @@ +"""Unit tests for label_transform — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-cpu-only") + +import unittest + +from sglang.srt.observability.label_transform import ( + UNKNOWN_PRIORITY_VALUE, + transform_priority, +) + + +class TestTransformPriority(unittest.TestCase): + """Test cases for transform_priority.""" + + def test_none_returns_unknown(self): + """None priority returns UNKNOWN.""" + self.assertEqual(transform_priority(None), UNKNOWN_PRIORITY_VALUE) + + def test_negative_returns_low(self): + """Priority below minimum returns LOW.""" + self.assertEqual(transform_priority(-1), "LOW") + + def test_above_max_returns_high(self): + """Priority at or above max returns HIGH.""" + self.assertEqual(transform_priority(31), "HIGH") + self.assertEqual(transform_priority(100), "HIGH") + + def test_in_range_returns_string(self): + """Priority in valid range [0, 31) returns its string representation.""" + self.assertEqual(transform_priority(0), "0") + self.assertEqual(transform_priority(15), "15") + self.assertEqual(transform_priority(30), "30") + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/observability/test_req_time_stats.py b/test/registered/unit/observability/test_req_time_stats.py new file mode 100644 index 000000000..d3de58695 --- /dev/null +++ b/test/registered/unit/observability/test_req_time_stats.py @@ -0,0 +1,837 @@ +"""Unit tests for req_time_stats.py — no server, no model loading.""" + +# ── Lightweight stubs for heavy transitive deps (torch, distributed, etc.) ── +# req_time_stats.py imports from disaggregation, model_executor, and other +# modules that transitively pull in torch/CUDA. We pre-populate sys.modules +# with minimal stubs so the module loads in a CPU-only environment. +import os +import sys +import types +from dataclasses import dataclass +from enum import Enum, IntEnum, auto + + +def _ensure_module(name): + if name not in sys.modules: + sys.modules[name] = types.ModuleType(name) + return sys.modules[name] + + +# Mirrors sglang.srt.disaggregation.utils.DisaggregationMode (stubbed to avoid torch dep). +class _DisaggregationMode(Enum): + NULL = "null" + PREFILL = "prefill" + DECODE = "decode" + + +def _kv_to_page_num(n, page_size): + return (n + page_size - 1) // page_size + + +_ensure_module("sglang.srt.disaggregation") +_du = _ensure_module("sglang.srt.disaggregation.utils") +_du.DisaggregationMode = _DisaggregationMode +_du.kv_to_page_num = _kv_to_page_num + + +# Mirrors sglang.srt.model_executor.forward_batch_info.ForwardMode (stubbed to avoid torch dep). +class _ForwardMode(IntEnum): + EXTEND = auto() + DECODE = auto() + MIXED = auto() + IDLE = auto() + TARGET_VERIFY = auto() + DRAFT_EXTEND = auto() + DRAFT_EXTEND_V2 = auto() + PREBUILT = auto() + SPLIT_PREFILL = auto() + DLLM_EXTEND = auto() + + def is_decode(self): + return self == _ForwardMode.DECODE + + def is_prefill(self): + return self in ( + _ForwardMode.EXTEND, + _ForwardMode.MIXED, + _ForwardMode.DRAFT_EXTEND, + _ForwardMode.TARGET_VERIFY, + _ForwardMode.SPLIT_PREFILL, + ) + + def is_prebuilt(self): + return self == _ForwardMode.PREBUILT + + +_ensure_module("sglang.srt.model_executor") +_fbi = _ensure_module("sglang.srt.model_executor.forward_batch_info") +_fbi.ForwardMode = _ForwardMode + + +# -- sglang.srt.observability.metrics_collector -- +_mc = _ensure_module("sglang.srt.observability.metrics_collector") +_mc.SchedulerMetricsCollector = type("SchedulerMetricsCollector", (), {}) +_mc.TokenizerMetricsCollector = type("TokenizerMetricsCollector", (), {}) + + +# -- sglang.srt.observability.trace -- +@dataclass +class _TraceNullContext: + tracing_enable: bool = False + + def __getattr__(self, name): + return self + + def __call__(self, *args, **kwargs): + return self + + +# -- sglang.srt.utils -- +# Stub both get_bool_env_var (for req_time_stats) and get_int_env_var (for trace.py) +# so the real trace module can load without torch. +def _get_bool_env_var(name, default="false"): + return os.getenv(name, default).lower() in ("true", "1") + + +def _get_int_env_var(name, default=0): + return int(os.getenv(name, str(default))) + + +_su = _ensure_module("sglang.srt.utils") +_su.get_bool_env_var = _get_bool_env_var +if not hasattr(_su, "get_int_env_var"): + _su.get_int_env_var = _get_int_env_var +_ensure_module("sglang.srt.utils.common") + +# ── End stubs ──────────────────────────────────────────────────────── + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import unittest +from unittest.mock import MagicMock + +import sglang.srt.observability.req_time_stats as rts_module +import sglang.srt.observability.trace as trace_module +from sglang.srt.observability.req_time_stats import ( + APIServerReqTimeStats, + DPControllerReqTimeStats, + ReqTimeStatsBase, + RequestStage, + RequestStageConfig, + SchedulerReqTimeStats, + calibrate_time_diff, + convert_time_cross_thread, + convert_time_to_realtime, + convert_time_to_realtime_ns, + monotonic_time, + real_time, + set_schedule_time_batch, + set_time_batch, +) +from sglang.srt.observability.trace import SpanAttributes + +DisaggregationMode = _DisaggregationMode +ForwardMode = _ForwardMode + + +class TestUtilityFunctions(unittest.TestCase): + def test_real_time_and_monotonic_time(self): + self.assertGreater(real_time(), 0) + self.assertGreater(monotonic_time(), 0) + + def test_convert_time_to_realtime(self): + result = convert_time_to_realtime(100.0) + self.assertIsInstance(result, float) + + def test_convert_time_to_realtime_ns(self): + result = convert_time_to_realtime_ns(100.0) + self.assertIsInstance(result, int) + self.assertGreater(result, 0) + + def test_convert_time_cross_thread(self): + self.assertAlmostEqual( + convert_time_cross_thread(10.0, old_diff=5.0, new_diff=3.0), 12.0 + ) + + def test_calibrate_time_diff(self): + calibrate_time_diff() # should not raise + + +class TestRequestStage(unittest.TestCase): + def test_stage_config_defaults(self): + cfg = RequestStageConfig("test_stage") + self.assertEqual(cfg.stage_name, "test_stage") + self.assertEqual(cfg.level, 0) + self.assertFalse(cfg.metrics_is_observed) + + def test_predefined_stages(self): + self.assertEqual(RequestStage.TOKENIZE.stage_name, "tokenize") + self.assertTrue(RequestStage.PREFILL_FORWARD.metrics_is_observed) + self.assertFalse(RequestStage.PREFILL_WAITING.metrics_is_observed) + self.assertEqual(RequestStage.ANONYMOUS.stage_name, "") + + +class TestReqTimeStatsBase(unittest.TestCase): + def test_disagg_mode_str(self): + base = ReqTimeStatsBase() + base.disagg_mode = DisaggregationMode.NULL + self.assertEqual(base.disagg_mode_str(), "unified") + base.disagg_mode = DisaggregationMode.DECODE + self.assertEqual(base.disagg_mode_str(), "decode") + base.disagg_mode = DisaggregationMode.PREFILL + self.assertEqual(base.disagg_mode_str(), "prefill") + + def test_set_metrics_collector(self): + base = ReqTimeStatsBase() + collector = MagicMock() + base.set_metrics_collector(collector) + self.assertTrue(base.enable_metrics) + self.assertIs(base.metrics_collector, collector) + + def test_set_metrics_collector_falsy(self): + base = ReqTimeStatsBase() + base.set_metrics_collector(None) + self.assertFalse(base.enable_metrics) + + def test_observe_per_stage_req_latency(self): + base = ReqTimeStatsBase() + collector = MagicMock() + base.set_metrics_collector(collector) + + stage = RequestStageConfig("test", metrics_is_observed=True) + base.observe_per_stage_req_latency(stage, 1.5) + collector.observe_per_stage_req_latency.assert_called_once_with("test", 1.5) + + def test_observe_per_stage_not_observed(self): + base = ReqTimeStatsBase() + collector = MagicMock() + base.set_metrics_collector(collector) + + stage = RequestStageConfig("test", metrics_is_observed=False) + base.observe_per_stage_req_latency(stage, 1.0) + collector.observe_per_stage_req_latency.assert_not_called() + + def test_init_trace_ctx(self): + base = ReqTimeStatsBase() + base.init_trace_ctx("rid-1", bootstrap_room=None) + # TraceReqContext stub has tracing_enable=False → replaced by TraceNullContext + self.assertFalse(base.trace_ctx.tracing_enable) + + def test_trace_slice_noop_when_tracing_disabled(self): + base = ReqTimeStatsBase() + base.trace_slice(RequestStage.TOKENIZE, 0.0, 1.0) + + def test_trace_slice_when_tracing_enabled(self): + base = ReqTimeStatsBase() + base.trace_ctx = MagicMock() + base.trace_ctx.tracing_enable = True + base.trace_slice(RequestStage.TOKENIZE, 0.0, 1.0, {"key": "val"}) + base.trace_ctx.trace_slice.assert_called_once() + + def test_new_from_obj_none(self): + obj = ReqTimeStatsBase.new_from_obj(None) + self.assertIsInstance(obj, ReqTimeStatsBase) + + def test_new_from_obj_copy(self): + src = ReqTimeStatsBase() + src.disagg_mode = DisaggregationMode.PREFILL + src.enable_metrics = True + dst = ReqTimeStatsBase.new_from_obj(src) + self.assertEqual(dst.disagg_mode, DisaggregationMode.PREFILL) + self.assertTrue(dst.enable_metrics) + + def test_getstate(self): + base = ReqTimeStatsBase() + base.disagg_mode = DisaggregationMode.DECODE + state = base.__getstate__() + self.assertEqual(state["disagg_mode"], DisaggregationMode.DECODE) + self.assertFalse(state["enable_metrics"]) + + def test_setstate_converts_time_fields(self): + base = ReqTimeStatsBase() + state = { + "created_time": 100.0, + "diff_realtime_monotonic": 50.0, + "disagg_mode": DisaggregationMode.NULL, + } + base.__setstate__(state) + self.assertEqual(base.disagg_mode, DisaggregationMode.NULL) + # created_time ends with "time" → converted via convert_time_cross_thread + expected = convert_time_cross_thread( + 100.0, 50.0, rts_module.global_diff_realtime_monotonic + ) + self.assertAlmostEqual(base.created_time, expected) + + +class TestAPIServerReqTimeStats(unittest.TestCase): + def test_setters_auto_timestamp(self): + """All setters work with default ts=None (auto perf_counter).""" + s = APIServerReqTimeStats() + s.set_created_time() + s.set_tokenize_finish_time() + s.set_api_server_dispatch_time() + s.set_api_server_dispatch_finish_time() + s.set_first_token_time() + s.set_last_time() + s.set_response_sent_to_client_time() + s.set_finished_time() + self.assertGreater(s.created_time, 0) + self.assertGreater(s.finished_time, 0) + + def test_getters(self): + s = APIServerReqTimeStats() + s.set_created_time(1.0) + s.set_first_token_time(3.0) + s.set_finished_time(5.0) + self.assertAlmostEqual(s.get_first_token_latency(), 2.0) + self.assertAlmostEqual(s.get_e2e_latency(), 4.0) + self.assertAlmostEqual(s.get_decode_latency(), 2.0) + + def test_get_interval(self): + s = APIServerReqTimeStats() + s.set_first_token_time(monotonic_time()) + self.assertGreaterEqual(s.get_interval(), 0.0) + + def test_getstate(self): + s = APIServerReqTimeStats(disagg_mode=DisaggregationMode.NULL) + state = s.__getstate__() + self.assertIn("disagg_mode", state) + self.assertFalse(state["enable_metrics"]) + + def test_get_response_sent_to_client_realtime(self): + s = APIServerReqTimeStats() + s.set_response_sent_to_client_time(100.0) + result = s.get_response_sent_to_client_realtime() + self.assertIsInstance(result, float) + + def test_convert_to_output_meta_info(self): + s = APIServerReqTimeStats() + s.set_created_time(1.0) + s.set_api_server_dispatch_finish_time(2.0) + s.set_first_token_time(3.0) + s.set_response_sent_to_client_time(4.0) + s.set_finished_time(5.0) + + meta = s.convert_to_output_meta_info(completion_tokens=10) + self.assertIn("request_received_ts", meta) + self.assertIn("api_server_dispatch_finish_ts", meta) + self.assertIn("response_sent_to_client_ts", meta) + self.assertIn("request_finished_ts", meta) + self.assertIn("decode_throughput", meta) + + def test_convert_to_output_meta_info_with_scheduler_stats(self): + s = APIServerReqTimeStats() + s.set_created_time(1.0) + s.set_first_token_time(3.0) + s.set_finished_time(5.0) + + sched = MagicMock() + sched.forward_entry_time = 2.0 + meta = s.convert_to_output_meta_info( + scheduler_time_stats=sched, completion_tokens=10 + ) + self.assertIn("inference_time", meta) + self.assertAlmostEqual(meta["inference_time"], 3.0) + + def test_convert_to_output_meta_info_empty(self): + """No timestamps set → minimal meta_info.""" + s = APIServerReqTimeStats() + meta = s.convert_to_output_meta_info() + self.assertNotIn("request_received_ts", meta) + self.assertNotIn("decode_throughput", meta) + + def test_convert_to_gen_ai_span_attrs(self): + s = APIServerReqTimeStats() + s.set_created_time(1.0) + s.set_first_token_time(3.0) + s.set_api_server_dispatch_finish_time(2.0) + s.set_finished_time(5.0) + + attrs = s.convert_to_gen_ai_span_attrs() + self.assertAlmostEqual( + attrs[SpanAttributes.GEN_AI_LATENCY_TIME_TO_FIRST_TOKEN], 2.0 + ) + self.assertAlmostEqual(attrs[SpanAttributes.GEN_AI_LATENCY_E2E], 4.0) + self.assertAlmostEqual( + attrs[SpanAttributes.GEN_AI_LATENCY_TIME_IN_MODEL_DECODE], 2.0 + ) + self.assertAlmostEqual( + attrs[SpanAttributes.GEN_AI_LATENCY_TIME_IN_MODEL_INFERENCE], 3.0 + ) + self.assertAlmostEqual( + attrs[SpanAttributes.GEN_AI_LATENCY_TIME_IN_MODEL_PREFILL], 1.0 + ) + + def test_convert_to_gen_ai_span_attrs_empty(self): + s = APIServerReqTimeStats() + attrs = s.convert_to_gen_ai_span_attrs() + self.assertEqual(len(attrs), 0) + + +class TestDPControllerReqTimeStats(unittest.TestCase): + def test_setters(self): + s = DPControllerReqTimeStats() + s.set_dp_dispatch_time() + s.set_dp_dispatch_finish_time() + self.assertGreater(s.dc_dispatch_time, 0) + self.assertGreater(s.dc_dispatch_finish_time, 0) + + def test_setters_explicit(self): + s = DPControllerReqTimeStats() + s.set_dp_dispatch_time(10.0) + s.set_dp_dispatch_finish_time(20.0) + self.assertEqual(s.dc_dispatch_time, 10.0) + self.assertEqual(s.dc_dispatch_finish_time, 20.0) + + def test_getstate(self): + s = DPControllerReqTimeStats(disagg_mode=DisaggregationMode.NULL) + state = s.__getstate__() + self.assertIn("disagg_mode", state) + + +class TestSchedulerReqTimeStats(unittest.TestCase): + def _make_stats(self, **kwargs): + defaults = dict(disagg_mode=DisaggregationMode.NULL) + defaults.update(kwargs) + return SchedulerReqTimeStats(**defaults) + + def _make_enabled_stats(self, **kwargs): + s = self._make_stats(**kwargs) + s.set_metrics_collector(MagicMock()) + return s + + def test_getstate_metrics_disabled(self): + s = self._make_stats() + self.assertEqual(s.__getstate__(), {}) + + def test_getstate_metrics_enabled(self): + s = self._make_enabled_stats() + s.wait_queue_entry_time = 1.0 + s.forward_entry_time = 2.0 + state = s.__getstate__() + self.assertIn("wait_queue_entry_time", state) + self.assertIn("forward_entry_time", state) + + def test_set_scheduler_recv_time(self): + s = self._make_stats() + s.set_scheduler_recv_time() + self.assertGreater(s.scheduler_recv_time, 0) + + def test_set_prefill_run_batch_times(self): + s = self._make_stats() + s.set_prefill_run_batch_start_time(1.0) + s.set_prefill_run_batch_end_time(2.0) + self.assertEqual(s.prefill_run_batch_start_time, 1.0) + self.assertEqual(s.prefill_run_batch_end_time, 2.0) + + def test_set_quick_finish_time(self): + s = self._make_stats() + s.set_quick_finish_time(5.0) + self.assertEqual(s.completion_time, 5.0) + self.assertEqual(s.forward_entry_time, 5.0) + + def test_set_bootstrap_done_time(self): + s = self._make_stats() + s.set_bootstrap_done_time(1.0) + self.assertEqual(s.bootstrap_done_time, 1.0) + # Second call does not overwrite + s.set_bootstrap_done_time(2.0) + self.assertEqual(s.bootstrap_done_time, 1.0) + + def test_set_completion_time(self): + s = self._make_stats() + s.set_completion_time(10.0) + self.assertEqual(s.completion_time, 10.0) + + def test_set_prefill_transfer_queue_entry_time(self): + s = self._make_stats() + s.set_prefill_transfer_queue_entry_time(1.0) + self.assertEqual(s.prefill_transfer_queue_entry_time, 1.0) + + def test_set_prefill_kv_transfer_finish_time(self): + s = self._make_stats() + s.prefill_transfer_queue_entry_time = 1.0 + s.set_prefill_kv_transfer_finish_time(3.0) + self.assertEqual(s.prefill_kv_transfer_finish_time, 3.0) + + def test_set_decode_prealloc_queue_entry_time(self): + s = self._make_stats() + s.scheduler_recv_time = 1.0 + s.set_decode_prealloc_queue_entry_time(2.0) + self.assertEqual(s.decode_prealloc_queue_entry_time, 2.0) + + def test_set_decode_transfer_queue_entry_time(self): + s = self._make_stats() + s.decode_prealloc_queue_entry_time = 1.0 + s.set_decode_transfer_queue_entry_time(2.0) + self.assertEqual(s.decode_transfer_queue_entry_time, 2.0) + + def test_set_decode_prebuilt_finish_time(self): + s = self._make_stats() + s.last_forward_entry_time = 1.0 + s.set_decode_prebuilt_finish_time(3.0) + self.assertEqual(s.decode_prebuilt_finish_time, 3.0) + + def test_set_prefill_bootstrap_queue_entry_time(self): + s = self._make_stats() + s.scheduler_recv_time = 1.0 + s.set_prefill_bootstrap_queue_entry_time(2.0) + self.assertEqual(s.prefill_bootstrap_queue_entry_time, 2.0) + + def test_auto_timestamp_setters(self): + """All setters default to perf_counter() when called without args.""" + s = self._make_stats() + s.set_scheduler_recv_time() + s.set_prefill_run_batch_start_time() + s.set_prefill_run_batch_end_time() + s.set_prefill_transfer_queue_entry_time() + s.set_prefill_kv_transfer_finish_time() + s.set_decode_prealloc_queue_entry_time() + s.set_decode_transfer_queue_entry_time() + s.set_bootstrap_done_time() + s.set_decode_prebuilt_finish_time() + s.set_completion_time() + s.set_quick_finish_time() + s.set_retract_time() + s.set_wait_queue_entry_time() + s.set_forward_entry_time() + s.set_last_chunked_prefill_finish_time() + s.set_prefill_finished_time() + s.set_last_decode_finish_time() + s.set_prefill_bootstrap_queue_entry_time() + s.set_last_scheduled_time(ForwardMode.DECODE) + self.assertGreater(s.scheduler_recv_time, 0) + self.assertGreater(s.completion_time, 0) + + def test_set_retract_time(self): + s = self._make_stats() + s.last_forward_entry_time = 1.0 + s.last_prefill_finished_time = 2.0 + s.set_retract_time(5.0) + self.assertEqual(s.last_forward_entry_time, 0.0) + self.assertEqual(s.last_prefill_finished_time, 0.0) + + def test_set_wait_queue_entry_time_first_call_null(self): + s = self._make_enabled_stats() + s.scheduler_recv_time = 1.0 + s.set_wait_queue_entry_time(3.0) + self.assertEqual(s.wait_queue_entry_time, 3.0) + s.metrics_collector.observe_per_stage_req_latency.assert_called() + + def test_set_wait_queue_entry_time_first_call_prefill(self): + s = self._make_enabled_stats(disagg_mode=DisaggregationMode.PREFILL) + s.prefill_bootstrap_queue_entry_time = 1.0 + s.set_wait_queue_entry_time(3.0) + self.assertEqual(s.wait_queue_entry_time, 3.0) + + def test_set_wait_queue_entry_time_first_call_decode(self): + s = self._make_enabled_stats(disagg_mode=DisaggregationMode.DECODE) + s.decode_transfer_queue_entry_time = 1.0 + s.set_wait_queue_entry_time(3.0) + self.assertEqual(s.wait_queue_entry_time, 3.0) + + def test_set_wait_queue_entry_time_retract(self): + s = self._make_stats() + s.wait_queue_entry_time = 1.0 # already set + s.set_wait_queue_entry_time(5.0) + self.assertEqual(s.wait_queue_entry_time, 5.0) + # retract resets these + self.assertEqual(s.last_forward_entry_time, 0.0) + + def test_set_forward_entry_time_first_call(self): + s = self._make_enabled_stats() + s.wait_queue_entry_time = 1.0 + s.set_forward_entry_time(3.0) + self.assertEqual(s.forward_entry_time, 3.0) + self.assertEqual(s.last_forward_entry_time, 3.0) + s.metrics_collector.observe_queue_time.assert_called_once() + + def test_set_forward_entry_time_first_call_decode(self): + s = self._make_enabled_stats(disagg_mode=DisaggregationMode.DECODE) + s.wait_queue_entry_time = 1.0 + s.set_forward_entry_time(3.0) + self.assertEqual(s.forward_entry_time, 3.0) + + def test_set_forward_entry_time_retract(self): + s = self._make_stats() + s.forward_entry_time = 1.0 # already set + s.last_forward_entry_time = 0.0 # reset by retract + s.set_forward_entry_time(5.0) + self.assertEqual(s.last_forward_entry_time, 5.0) + + def test_set_last_chunked_prefill_finish_time_first(self): + s = self._make_stats() + s.last_forward_entry_time = 1.0 + s.set_last_chunked_prefill_finish_time(3.0) + self.assertEqual(s.last_chunked_prefill_finish_time, 3.0) + + def test_set_last_chunked_prefill_finish_time_subsequent(self): + s = self._make_stats() + s.last_forward_entry_time = 1.0 + s.last_chunked_prefill_finish_time = 2.0 + s.set_last_chunked_prefill_finish_time(4.0) + self.assertEqual(s.last_chunked_prefill_finish_time, 4.0) + + def test_set_prefill_finished_time_first_call(self): + s = self._make_enabled_stats() + s.last_forward_entry_time = 1.0 + s.set_prefill_finished_time(3.0) + self.assertEqual(s.prefill_finished_time, 3.0) + self.assertEqual(s.last_prefill_finished_time, 3.0) + + def test_set_prefill_finished_time_retract(self): + s = self._make_stats() + s.prefill_finished_time = 1.0 # already set + s.last_prefill_finished_time = 0.0 # reset by retract + s.last_forward_entry_time = 0.5 + s.set_prefill_finished_time(5.0) + self.assertEqual(s.last_prefill_finished_time, 5.0) + + def test_set_prefill_finished_time_retract_with_chunked(self): + s = self._make_stats() + s.prefill_finished_time = 1.0 + s.last_prefill_finished_time = 0.0 + s.last_chunked_prefill_finish_time = 2.0 + s.set_prefill_finished_time(5.0) + self.assertEqual(s.last_prefill_finished_time, 5.0) + + def test_set_prefill_finished_time_tracing_enabled(self): + s = self._make_enabled_stats() + s.trace_ctx = MagicMock() + s.trace_ctx.tracing_enable = True + s.last_forward_entry_time = 1.0 + s.last_chunked_prefill_finish_time = 2.0 + s.last_decode_scheduled_time = 1.5 + s.set_prefill_finished_time(3.0) + self.assertEqual(s.prefill_finished_time, 3.0) + s.trace_ctx.trace_slice_end.assert_called_once() + # NULL mode + last_decode_scheduled_time > 0 → trace_slice_start for decode + s.trace_ctx.trace_slice_start.assert_called_once() + + def test_set_last_decode_finish_time_first(self): + s = self._make_enabled_stats() + s.last_prefill_finished_time = 1.0 + s.last_decode_scheduled_time = 0.5 + s.set_last_decode_finish_time(3.0) + self.assertEqual(s.last_decode_finish_time, 3.0) + self.assertEqual(s.decode_ct, 1) + + def test_set_last_decode_finish_time_first_decode_mode(self): + s = self._make_enabled_stats(disagg_mode=DisaggregationMode.DECODE) + s.decode_prebuilt_finish_time = 1.0 + s.set_last_decode_finish_time(3.0) + self.assertEqual(s.decode_ct, 1) + + def test_set_last_decode_finish_time_first_with_scheduled(self): + s = self._make_enabled_stats() + s.last_prefill_finished_time = 1.0 + s.last_decode_scheduled_time = 2.0 + s.set_last_decode_finish_time(3.0) + self.assertEqual(s.decode_ct, 1) + + def test_set_last_decode_finish_time_subsequent(self): + s = self._make_enabled_stats() + s.last_decode_finish_time = 1.0 + s.set_last_decode_finish_time(3.0) + self.assertEqual(s.decode_ct, 1) + + def test_set_last_scheduled_time_decode(self): + s = self._make_stats() + s.set_last_scheduled_time(ForwardMode.DECODE, 5.0) + self.assertEqual(s.last_decode_scheduled_time, 5.0) + + def test_set_last_scheduled_time_extend(self): + s = self._make_stats() + s.set_last_scheduled_time(ForwardMode.EXTEND, 5.0) + self.assertEqual(s.last_decode_scheduled_time, 0.0) # not decode + + def test_set_last_scheduled_time_tracing_enabled(self): + s = self._make_stats() + s.trace_ctx = MagicMock() + s.trace_ctx.tracing_enable = True + s.last_prefill_finished_time = 1.0 + s.set_last_scheduled_time(ForwardMode.DECODE, 5.0) + s.trace_ctx.trace_event.assert_called_once() + # NULL mode + first decode + last_prefill_finished_time > 0 → trace decode waiting + self.assertEqual(s.last_decode_scheduled_time, 5.0) + + def test_get_queueing_time(self): + s = self._make_stats() + s.forward_entry_time = 5.0 + s.wait_queue_entry_time = 2.0 + self.assertAlmostEqual(s.get_queueing_time(), 3.0) + + def test_get_prefill_waiting_latency(self): + s = self._make_stats() + self.assertIsNone(s.get_prefill_waiting_latency()) + s.prefill_run_batch_start_time = 3.0 + s.forward_entry_time = 1.0 + self.assertAlmostEqual(s.get_prefill_waiting_latency(), 2.0) + + def test_get_prefill_launch_latency(self): + s = self._make_stats() + self.assertIsNone(s.get_prefill_launch_latency()) + s.prefill_run_batch_start_time = 1.0 + s.prefill_run_batch_end_time = 3.0 + self.assertAlmostEqual(s.get_prefill_launch_latency(), 2.0) + + def test_format_duration(self): + s = self._make_stats() + self.assertEqual(s.format_duration(0.001), "1.00ms") + + def test_convert_to_duration_null(self): + s = self._make_stats( + wait_queue_entry_time=1.0, + forward_entry_time=2.0, + completion_time=5.0, + ) + result = s.convert_to_duration() + self.assertIn("queue_duration=", result) + self.assertIn("forward_duration=", result) + + def test_convert_to_duration_prefill_no_bootstrap(self): + s = self._make_stats( + disagg_mode=DisaggregationMode.PREFILL, + prefill_bootstrap_queue_entry_time=1.0, + wait_queue_entry_time=2.0, + forward_entry_time=3.0, + completion_time=5.0, + ) + result = s.convert_to_duration() + self.assertIn("bootstrap_queue_duration", result) + self.assertNotIn("alloc_wait", result) + + def test_convert_to_duration_prefill_with_bootstrap(self): + s = self._make_stats( + disagg_mode=DisaggregationMode.PREFILL, + prefill_bootstrap_queue_entry_time=1.0, + bootstrap_done_time=1.5, + wait_queue_entry_time=2.0, + forward_entry_time=3.0, + completion_time=5.0, + ) + result = s.convert_to_duration() + self.assertIn("bootstrap(", result) + self.assertIn("alloc_wait(", result) + + def test_convert_to_duration_decode_no_bootstrap(self): + s = self._make_stats( + disagg_mode=DisaggregationMode.DECODE, + decode_prealloc_queue_entry_time=1.0, + decode_transfer_queue_entry_time=2.0, + wait_queue_entry_time=3.0, + forward_entry_time=4.0, + completion_time=6.0, + ) + result = s.convert_to_duration() + self.assertIn("prealloc_queue_duration", result) + self.assertIn("transfer_duration=", result) + self.assertNotIn("alloc_wait", result) + + def test_convert_to_duration_decode_with_bootstrap(self): + s = self._make_stats( + disagg_mode=DisaggregationMode.DECODE, + decode_prealloc_queue_entry_time=1.0, + bootstrap_done_time=1.5, + decode_transfer_queue_entry_time=2.0, + wait_queue_entry_time=3.0, + forward_entry_time=4.0, + completion_time=6.0, + ) + result = s.convert_to_duration() + self.assertIn("bootstrap(", result) + self.assertIn("alloc_wait(", result) + + def test_convert_to_duration_unknown_mode(self): + s = self._make_stats() + # Force an invalid disagg_mode to hit the else branch + s.disagg_mode = "invalid" + self.assertEqual(s.convert_to_duration(), "Unknown Time Stats") + + def test_convert_to_output_meta_info(self): + s = self._make_stats( + forward_entry_time=2.0, + prefill_finished_time=3.0, + wait_queue_entry_time=1.0, + prefill_run_batch_start_time=2.5, + prefill_run_batch_end_time=2.8, + ) + meta = s.convert_to_output_meta_info() + self.assertIn("forward_entry_time", meta) + self.assertIn("prefill_finished_time", meta) + self.assertIn("queue_time", meta) + self.assertIn("prefill_waiting_latency", meta) + self.assertIn("prefill_launch_latency", meta) + + def test_compute_kv_transfer_metrics(self): + s = self._make_enabled_stats() + s.prefill_transfer_queue_entry_time = 1.0 + s.completion_time = 2.0 + result = s.compute_and_observe_kv_transfer_metrics( + num_tokens=10, page_size=4, bytes_per_page_all_layers=1024 + ) + self.assertIn("latency_ms", result) + self.assertIn("total_mb", result) + self.assertIn("speed_gb_s", result) + s.metrics_collector.observe_kv_transfer_metrics.assert_called_once() + + def test_compute_kv_transfer_with_bootstrap(self): + s = self._make_enabled_stats() + s.prefill_transfer_queue_entry_time = 1.0 + s.completion_time = 2.0 + s.prefill_bootstrap_queue_entry_time = 0.5 + s.bootstrap_done_time = 0.8 + s.wait_queue_entry_time = 1.0 + result = s.compute_and_observe_kv_transfer_metrics( + num_tokens=10, page_size=4, bytes_per_page_all_layers=1024 + ) + self.assertIn("bootstrap_ms", result) + self.assertIn("alloc_ms", result) + s.metrics_collector.observe_kv_transfer_bootstrap.assert_called_once() + + def test_compute_kv_transfer_metrics_none(self): + s = self._make_stats() + result = s.compute_and_observe_kv_transfer_metrics( + num_tokens=10, page_size=4, bytes_per_page_all_layers=1024 + ) + self.assertIsNone(result) + + +class TestBatchFunctions(unittest.TestCase): + def test_set_time_batch_empty(self): + set_time_batch(None, "set_completion_time") + set_time_batch([], "set_completion_time") + + def test_set_time_batch(self): + req = MagicMock() + set_time_batch([req], "set_completion_time") + req.time_stats.set_completion_time.assert_called_once() + + def test_set_schedule_time_batch_tracing_disabled(self): + batch = MagicMock() + set_schedule_time_batch(batch) + # Tracing is disabled in our stub → returns early + + def test_set_schedule_time_batch_tracing_enabled(self): + orig = trace_module.get_global_tracing_enabled + trace_module.get_global_tracing_enabled = lambda: True + rts_module.get_global_tracing_enabled = lambda: True + try: + req = MagicMock() + batch = MagicMock() + batch.reqs = [req] + batch.forward_mode = ForwardMode.DECODE + batch.forward_mode.is_decode = lambda: True + batch.forward_mode.is_prefill = lambda: False + batch.forward_mode.is_prebuilt = lambda: False + set_schedule_time_batch(batch) + req.time_stats.set_last_scheduled_time.assert_called_once() + finally: + trace_module.get_global_tracing_enabled = orig + rts_module.get_global_tracing_enabled = orig + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/observability/test_request_metrics_exporter.py b/test/registered/unit/observability/test_request_metrics_exporter.py new file mode 100644 index 000000000..f808f3f7c --- /dev/null +++ b/test/registered/unit/observability/test_request_metrics_exporter.py @@ -0,0 +1,323 @@ +"""Unit tests for request_metrics_exporter.py — no server, no model loading.""" + +# ── Lightweight stubs for heavy transitive deps ── +import sys +import types +from dataclasses import dataclass +from typing import Any, Dict, List, Optional + + +def _ensure_module(name): + if name not in sys.modules: + sys.modules[name] = types.ModuleType(name) + return sys.modules[name] + + +@dataclass +class _GenerateReqInput: + rid: Optional[str] = None + text: Optional[str] = None + image_data: Optional[Any] = None # in ALWAYS_EXCLUDE_FIELDS + sampling_params: Optional[Dict] = None + + +@dataclass +class _EmbeddingReqInput: + rid: Optional[str] = None + text: Optional[str] = None + image_data: Optional[Any] = None + input_ids: Optional[List[int]] = None + + +class _ServerArgs: + def __init__(self, **kwargs): + for k, v in kwargs.items(): + setattr(self, k, v) + + +# Pre-populate modules before importing the module under test. +_ensure_module("sglang.srt.managers") +_io = _ensure_module("sglang.srt.managers.io_struct") +_io.GenerateReqInput = _GenerateReqInput +_io.EmbeddingReqInput = _EmbeddingReqInput + +_sa = _ensure_module("sglang.srt.server_args") +_sa.ServerArgs = _ServerArgs + +# ── End stubs ── + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import asyncio +import json +import os +import shutil +import tempfile +import unittest +from unittest.mock import MagicMock, patch + +from sglang.srt.observability.request_metrics_exporter import ( + FileRequestMetricsExporter, + RequestMetricsExporter, + RequestMetricsExporterManager, + create_request_metrics_exporters, +) + + +def _make_server_args(tmp_dir, enabled=True): + return _ServerArgs( + export_metrics_to_file=enabled, + export_metrics_to_file_dir=tmp_dir, + ) + + +class _ConcreteExporter(RequestMetricsExporter): + """Minimal concrete subclass for testing base class methods.""" + + async def write_record(self, obj, out_dict): + pass + + +class TestFormatOutputData(unittest.TestCase): + def test_basic_formatting(self): + server_args = _make_server_args("/tmp/unused") + exporter = _ConcreteExporter( + server_args, obj_skip_names=None, out_skip_names=None + ) + + obj = _GenerateReqInput( + rid="req-1", text="hello", sampling_params={"temp": 0.5} + ) + out_dict = {"meta_info": {"latency": 1.5, "tokens": 10}} + + result = exporter._format_output_data(obj, out_dict) + + params = json.loads(result["request_parameters"]) + self.assertEqual(params["rid"], "req-1") + self.assertEqual(params["text"], "hello") + self.assertIn("latency", result) + self.assertIn("tokens", result) + + def test_excludes_always_exclude_fields(self): + server_args = _make_server_args("/tmp/unused") + exporter = _ConcreteExporter( + server_args, obj_skip_names=None, out_skip_names=None + ) + + obj = _GenerateReqInput(rid="req-1", image_data="should_be_excluded") + result = exporter._format_output_data(obj, {}) + + params = json.loads(result["request_parameters"]) + self.assertNotIn("image_data", params) + + def test_excludes_obj_skip_names(self): + server_args = _make_server_args("/tmp/unused") + exporter = _ConcreteExporter( + server_args, obj_skip_names={"text"}, out_skip_names=None + ) + + obj = _GenerateReqInput(rid="req-1", text="skip_me") + result = exporter._format_output_data(obj, {}) + + params = json.loads(result["request_parameters"]) + self.assertNotIn("text", params) + self.assertIn("rid", params) + + def test_excludes_none_values(self): + server_args = _make_server_args("/tmp/unused") + exporter = _ConcreteExporter( + server_args, obj_skip_names=None, out_skip_names=None + ) + + obj = _GenerateReqInput(rid="req-1", text=None) + result = exporter._format_output_data(obj, {}) + + params = json.loads(result["request_parameters"]) + self.assertNotIn("text", params) + + def test_filters_out_skip_names(self): + server_args = _make_server_args("/tmp/unused") + exporter = _ConcreteExporter( + server_args, obj_skip_names=None, out_skip_names={"secret"} + ) + + obj = _GenerateReqInput(rid="req-1") + out_dict = {"meta_info": {"latency": 1.5, "secret": "hidden"}} + result = exporter._format_output_data(obj, out_dict) + + self.assertIn("latency", result) + self.assertNotIn("secret", result) + + +class TestFileRequestMetricsExporter(unittest.TestCase): + def setUp(self): + self.tmp_dir = tempfile.mkdtemp() + + def tearDown(self): + shutil.rmtree(self.tmp_dir, ignore_errors=True) + + def _make_exporter(self): + return FileRequestMetricsExporter(_make_server_args(self.tmp_dir), None, None) + + def test_init_creates_directory(self): + sub_dir = os.path.join(self.tmp_dir, "nested", "dir") + FileRequestMetricsExporter(_make_server_args(sub_dir), None, None) + self.assertTrue(os.path.isdir(sub_dir)) + + def test_ensure_file_handler_opens_file(self): + exporter = self._make_exporter() + exporter._ensure_file_handler("20240101_12") + self.assertIsNotNone(exporter._current_file_handler) + self.assertEqual(exporter._current_hour_suffix, "20240101_12") + exporter.close() + + def test_ensure_file_handler_rotates(self): + exporter = self._make_exporter() + exporter._ensure_file_handler("20240101_12") + first_handler = exporter._current_file_handler + exporter._ensure_file_handler("20240101_13") + self.assertTrue(first_handler.closed) + self.assertEqual(exporter._current_hour_suffix, "20240101_13") + exporter.close() + + def test_ensure_file_handler_close_error(self): + """Previous handler close failure is logged but doesn't prevent rotation.""" + exporter = self._make_exporter() + mock_handler = MagicMock() + mock_handler.close.side_effect = OSError("disk error") + exporter._current_file_handler = mock_handler + exporter._current_hour_suffix = "old" + + exporter._ensure_file_handler("new") + self.assertEqual(exporter._current_hour_suffix, "new") + exporter.close() + + def test_ensure_file_handler_open_error(self): + exporter = self._make_exporter() + with patch("builtins.open", side_effect=OSError("permission denied")): + with self.assertRaises(OSError): + exporter._ensure_file_handler("20240101_12") + self.assertIsNone(exporter._current_file_handler) + self.assertIsNone(exporter._current_hour_suffix) + + def test_close(self): + exporter = self._make_exporter() + exporter._ensure_file_handler("20240101_12") + exporter.close() + self.assertIsNone(exporter._current_file_handler) + self.assertIsNone(exporter._current_hour_suffix) + + def test_close_noop_when_no_handler(self): + exporter = self._make_exporter() + exporter.close() # should not raise + + def test_close_error(self): + """Close failure is logged but state is still reset.""" + exporter = self._make_exporter() + mock_handler = MagicMock() + mock_handler.close.side_effect = OSError("disk error") + exporter._current_file_handler = mock_handler + exporter._current_hour_suffix = "old" + + exporter.close() + self.assertIsNone(exporter._current_file_handler) + self.assertIsNone(exporter._current_hour_suffix) + + def test_write_record(self): + exporter = self._make_exporter() + obj = _GenerateReqInput(rid="req-1", text="hello") + out_dict = {"meta_info": {"latency": 1.5}} + + asyncio.run(exporter.write_record(obj, out_dict)) + + # Find the written file + files = os.listdir(self.tmp_dir) + self.assertEqual(len(files), 1) + with open(os.path.join(self.tmp_dir, files[0])) as f: + record = json.loads(f.readline()) + self.assertIn("request_parameters", record) + self.assertAlmostEqual(record["latency"], 1.5) + exporter.close() + + def test_write_record_skips_health_check(self): + exporter = self._make_exporter() + obj = _GenerateReqInput(rid="HEALTH_CHECK_123", text="ping") + asyncio.run(exporter.write_record(obj, {})) + + files = os.listdir(self.tmp_dir) + self.assertEqual(len(files), 0) + + def test_write_record_handler_none(self): + """If file handler is None after ensure, write_record returns early.""" + exporter = self._make_exporter() + obj = _GenerateReqInput(rid="req-1") + + with patch.object(exporter, "_ensure_file_handler"): + exporter._current_file_handler = None + asyncio.run(exporter.write_record(obj, {})) + # No crash, no file written + + def test_write_record_exception(self): + """Exceptions during write are caught and logged.""" + exporter = self._make_exporter() + obj = _GenerateReqInput(rid="req-1") + + with patch.object( + exporter, "_ensure_file_handler", side_effect=RuntimeError("boom") + ): + asyncio.run(exporter.write_record(obj, {})) + # Should not raise + + +class TestRequestMetricsExporterManager(unittest.TestCase): + def setUp(self): + self.tmp_dir = tempfile.mkdtemp() + + def tearDown(self): + shutil.rmtree(self.tmp_dir, ignore_errors=True) + + def test_no_exporters(self): + server_args = _make_server_args(self.tmp_dir, enabled=False) + manager = RequestMetricsExporterManager(server_args) + self.assertFalse(manager.exporter_enabled()) + + def test_with_file_exporter(self): + server_args = _make_server_args(self.tmp_dir, enabled=True) + manager = RequestMetricsExporterManager(server_args) + self.assertTrue(manager.exporter_enabled()) + + def test_write_record_delegates(self): + server_args = _make_server_args(self.tmp_dir, enabled=True) + manager = RequestMetricsExporterManager(server_args) + + obj = _GenerateReqInput(rid="req-1", text="hello") + out_dict = {"meta_info": {"latency": 1.0}} + asyncio.run(manager.write_record(obj, out_dict)) + + files = os.listdir(self.tmp_dir) + self.assertEqual(len(files), 1) + + +class TestCreateExporters(unittest.TestCase): + def setUp(self): + self.tmp_dir = tempfile.mkdtemp() + + def tearDown(self): + shutil.rmtree(self.tmp_dir, ignore_errors=True) + + def test_disabled(self): + server_args = _make_server_args(self.tmp_dir, enabled=False) + exporters = create_request_metrics_exporters(server_args) + self.assertEqual(len(exporters), 0) + + def test_enabled(self): + server_args = _make_server_args(self.tmp_dir, enabled=True) + exporters = create_request_metrics_exporters(server_args) + self.assertEqual(len(exporters), 1) + self.assertIsInstance(exporters[0], FileRequestMetricsExporter) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/observability/test_startup_func_log_and_timer.py b/test/registered/unit/observability/test_startup_func_log_and_timer.py new file mode 100644 index 000000000..a5d46e22f --- /dev/null +++ b/test/registered/unit/observability/test_startup_func_log_and_timer.py @@ -0,0 +1,140 @@ +"""Unit tests for startup_func_log_and_timer.py — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import unittest +from unittest.mock import MagicMock, patch + +import sglang.srt.observability.startup_func_log_and_timer as mod +from sglang.srt.observability.startup_func_log_and_timer import ( + enable_startup_timer, + get_max_duration, + reset_startup_timers, + set_startup_metric, + startup_timer, + time_startup_latency, +) + + +class TestStartupFuncLogAndTimer(unittest.TestCase): + def setUp(self): + self.orig_enable = mod.enable_startup_metrics + self.orig_gauge = mod.STARTUP_LATENCY_SECONDS + mod._max_durations.clear() + + def tearDown(self): + mod.enable_startup_metrics = self.orig_enable + mod.STARTUP_LATENCY_SECONDS = self.orig_gauge + mod._max_durations.clear() + + @patch("prometheus_client.Gauge") + def test_enable_startup_timer(self, MockGauge): + enable_startup_timer() + self.assertTrue(mod.enable_startup_metrics) + self.assertIs(mod.STARTUP_LATENCY_SECONDS, MockGauge.return_value) + MockGauge.assert_called_once() + + def test_reset_and_get_max_duration(self): + mod._max_durations["ctx"] = 5.0 + self.assertAlmostEqual(get_max_duration("ctx"), 5.0) + self.assertIsNone(get_max_duration("nonexistent")) + reset_startup_timers() + self.assertIsNone(get_max_duration("ctx")) + + def test_set_startup_metric_disabled(self): + """When metrics disabled, returns early without tracking max.""" + mod.enable_startup_metrics = False + set_startup_metric("ctx", 1.0) + self.assertIsNone(get_max_duration("ctx")) + + def test_set_startup_metric_enabled(self): + """Tracks max and updates gauge when enabled.""" + mock_gauge = MagicMock() + mod.enable_startup_metrics = True + mod.STARTUP_LATENCY_SECONDS = mock_gauge + + set_startup_metric("ctx", 1.0) + self.assertAlmostEqual(get_max_duration("ctx"), 1.0) + mock_gauge.labels.assert_called_with(context="ctx") + + # Lower value → not updated + mock_gauge.reset_mock() + set_startup_metric("ctx", 0.5) + self.assertAlmostEqual(get_max_duration("ctx"), 1.0) + mock_gauge.labels().set.assert_not_called() + + def test_set_startup_metric_no_log(self): + mod.enable_startup_metrics = False + with patch.object(mod.logger, "info") as mock_log: + set_startup_metric("ctx", 1.0, should_log=False) + mock_log.assert_not_called() + + def test_startup_timer_basic(self): + with startup_timer("block"): + pass + self.assertGreaterEqual(get_max_duration("block"), 0.0) + + def test_startup_timer_with_gauge(self): + """Gauge updated when metrics enabled and log_only=False.""" + mock_gauge = MagicMock() + mod.enable_startup_metrics = True + mod.STARTUP_LATENCY_SECONDS = mock_gauge + + with startup_timer("block"): + pass + mock_gauge.labels.assert_called_with(context="block") + mock_gauge.labels().set.assert_called_once() + + def test_startup_timer_log_only(self): + """log_only=True skips gauge but still tracks max.""" + mock_gauge = MagicMock() + mod.enable_startup_metrics = True + mod.STARTUP_LATENCY_SECONDS = mock_gauge + + with startup_timer("block", log_only=True): + pass + mock_gauge.labels.assert_not_called() + self.assertIsNotNone(get_max_duration("block")) + + def test_decorator_direct(self): + """Direct decorator @time_startup_latency preserves return value.""" + + @time_startup_latency + def add(a, b): + return a + b + + self.assertEqual(add(2, 3), 5) + self.assertIsNotNone(get_max_duration("add")) + + def test_decorator_factory_with_gauge(self): + """Factory decorator with custom name, gauge updated.""" + mock_gauge = MagicMock() + mod.enable_startup_metrics = True + mod.STARTUP_LATENCY_SECONDS = mock_gauge + + @time_startup_latency(name="custom_op") + def add(a, b): + return a + b + + self.assertEqual(add(2, 3), 5) + mock_gauge.labels.assert_called_with(context="custom_op") + + def test_decorator_log_only(self): + """log_only=True skips gauge but still tracks max.""" + mock_gauge = MagicMock() + mod.enable_startup_metrics = True + mod.STARTUP_LATENCY_SECONDS = mock_gauge + + @time_startup_latency(log_only=True) + def add(a, b): + return a + b + + self.assertEqual(add(2, 3), 5) + mock_gauge.labels.assert_not_called() + self.assertIsNotNone(get_max_duration("add")) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/observability/test_trace.py b/test/registered/unit/observability/test_trace.py new file mode 100644 index 000000000..9e86579b8 --- /dev/null +++ b/test/registered/unit/observability/test_trace.py @@ -0,0 +1,605 @@ +"""Unit tests for trace.py — no server, no model loading.""" + +# ── Stubs for heavy transitive deps ── +import os +import sys +import types + + +def _ensure_module(name): + if name not in sys.modules: + sys.modules[name] = types.ModuleType(name) + return sys.modules[name] + + +_su = _ensure_module("sglang.srt.utils") +if not hasattr(_su, "get_int_env_var"): + _su.get_int_env_var = lambda name, default=0: int(os.getenv(name, str(default))) +_ensure_module("sglang.srt.utils.common") + +_ensure_module("sglang.srt.managers") +_sb = _ensure_module("sglang.srt.managers.schedule_batch") +_sb.BaseFinishReason = type("BaseFinishReason", (), {"to_json": lambda self: {}}) + +# ── End stubs ── + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import threading +import unittest +from unittest.mock import patch + +import sglang.srt.observability.trace as mod +from sglang.srt.observability.trace import ( + SpanAttributes, + TraceCustomIdGenerator, + TraceEvent, + TraceNullContext, + TraceReqContext, + TraceSliceContext, + TraceThreadContext, + TraceThreadInfo, + extract_trace_headers, + get_global_tracing_enabled, + process_tracing_init, + set_global_trace_level, + trace_set_thread_info, +) + +try: + from opentelemetry import trace as otel_trace + from opentelemetry.sdk.trace import TracerProvider + + from sglang.srt.observability.trace import get_otlp_span_exporter + + _has_otel = True +except ImportError: + _has_otel = False + +# Access the private module-level function (avoid name mangling inside classes). +_get_host_id = getattr(mod, "__get_host_id") + + +class TestTraceFunctions(unittest.TestCase): + def test_extract_trace_headers(self): + headers = {"traceparent": "abc", "tracestate": "xyz", "other": "skip"} + result = extract_trace_headers(headers) + self.assertEqual(result, {"traceparent": "abc", "tracestate": "xyz"}) + + def test_extract_trace_headers_missing(self): + self.assertEqual(extract_trace_headers({}), {}) + + def test_set_global_trace_level(self): + orig = mod.global_trace_level + set_global_trace_level(5) + self.assertEqual(mod.global_trace_level, 5) + mod.global_trace_level = orig + + def test_get_global_tracing_enabled(self): + self.assertEqual(get_global_tracing_enabled(), mod.opentelemetry_initialized) + + def test_get_cur_time_ns(self): + ts = mod.get_cur_time_ns() + self.assertIsInstance(ts, int) + self.assertGreater(ts, 0) + + +class TestDataclasses(unittest.TestCase): + def test_trace_thread_info(self): + info = TraceThreadInfo("host", 123, "label", 0, 1) + self.assertEqual(info.thread_label, "label") + + def test_trace_event(self): + evt = TraceEvent("name", 100, {"k": "v"}) + self.assertEqual(evt.event_name, "name") + + def test_trace_slice_context(self): + s = TraceSliceContext("slice", 100, end_time_ns=200, level=2, attrs={"a": 1}) + self.assertEqual(s.slice_name, "slice") + + def test_trace_thread_context(self): + info = TraceThreadInfo("h", 1, "l", 0, 0) + ctx = TraceThreadContext(thread_info=info, cur_slice_stack=[]) + self.assertEqual(len(ctx.cur_slice_stack), 0) + + +class TestTraceNullContext(unittest.TestCase): + def test_null_object_pattern(self): + ctx = TraceNullContext() + self.assertFalse(ctx.tracing_enable) + # Any attribute access returns self + self.assertIs(ctx.some_method, ctx) + # Callable returns self + self.assertIs(ctx("arg1", key="val"), ctx) + # Chaining works + self.assertIs(ctx.foo.bar.baz(1, 2, 3), ctx) + + +class TestSpanAttributes(unittest.TestCase): + def test_constants_exist(self): + self.assertEqual(SpanAttributes.GEN_AI_LATENCY_E2E, "gen_ai.latency.e2e") + self.assertIsInstance(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS, str) + + +class TestTraceCustomIdGenerator(unittest.TestCase): + def test_generates_nonzero_ids(self): + gen = TraceCustomIdGenerator() + trace_id = gen.generate_trace_id() + span_id = gen.generate_span_id() + self.assertIsInstance(trace_id, int) + self.assertIsInstance(span_id, int) + + +# __get_host_id +class TestGetHostId(unittest.TestCase): + def test_from_machine_id_file(self): + with patch("os.path.exists", return_value=True), patch( + "builtins.open", + unittest.mock.mock_open(read_data="abc123\n"), + ): + self.assertEqual(_get_host_id(), "abc123") + + def test_from_machine_id_file_error(self): + """Falls back to MAC address when file read fails.""" + with patch("os.path.exists", return_value=True), patch( + "builtins.open", side_effect=IOError("read error") + ): + result = _get_host_id() + self.assertIsInstance(result, str) + self.assertGreater(len(result), 0) + + def test_from_mac_address(self): + with patch("os.path.exists", return_value=False), patch( + "uuid.getnode", return_value=0x112233445566 + ): + result = _get_host_id() + self.assertIsInstance(result, str) + self.assertGreater(len(result), 0) + + def test_unknown_fallback(self): + with patch("os.path.exists", return_value=False), patch( + "uuid.getnode", return_value=0 + ): + self.assertEqual(_get_host_id(), "unknown") + + +@unittest.skipUnless(_has_otel, "opentelemetry not installed") +class TestGetOtlpSpanExporter(unittest.TestCase): + def test_grpc_default(self): + + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("OTEL_EXPORTER_OTLP_TRACES_PROTOCOL", None) + exporter = get_otlp_span_exporter("localhost:4317") + self.assertIsNotNone(exporter) + + def test_http_protobuf(self): + + with patch.dict( + os.environ, {"OTEL_EXPORTER_OTLP_TRACES_PROTOCOL": "http/protobuf"} + ): + exporter = get_otlp_span_exporter("http://localhost:4318/v1/traces") + self.assertIsNotNone(exporter) + + def test_invalid_protocol(self): + + with patch.dict(os.environ, {"OTEL_EXPORTER_OTLP_TRACES_PROTOCOL": "invalid"}): + with self.assertRaises(ValueError): + get_otlp_span_exporter("localhost:4317") + + +class TestProcessTracingInit(unittest.TestCase): + def test_raises_without_otel(self): + + orig = mod.opentelemetry_imported + mod.opentelemetry_imported = False + try: + with self.assertRaises(RuntimeError): + process_tracing_init("localhost:4317", "test") + finally: + mod.opentelemetry_imported = orig + + +class TestTraceReqContextDisabled(unittest.TestCase): + def setUp(self): + self.orig = mod.opentelemetry_initialized + mod.opentelemetry_initialized = False + + def tearDown(self): + mod.opentelemetry_initialized = self.orig + + def test_init_disabled(self): + ctx = TraceReqContext(rid="req-1") + self.assertFalse(ctx.tracing_enable) + self.assertFalse(ctx.is_tracing_enabled()) + + def test_all_methods_noop(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start() + ctx.trace_req_finish() + ctx.trace_slice_start("s", 1) + ctx.trace_slice_end("s", 1) + ctx.trace_slice(TraceSliceContext("s", 100)) + ctx.trace_event("e", 1) + ctx.trace_set_root_attrs({"k": "v"}) + ctx.trace_set_thread_attrs({"k": "v"}) + ctx.abort() + ctx.rebuild_thread_context() + + def test_getstate_disabled(self): + ctx = TraceReqContext(rid="req-1") + state = ctx.__getstate__() + self.assertEqual(state, {"tracing_enable": False}) + + def test_setstate_disabled(self): + ctx = TraceReqContext.__new__(TraceReqContext) + ctx.__setstate__({"tracing_enable": True, "is_copy": False}) + # opentelemetry_initialized is False → tracing forced off + self.assertFalse(ctx.tracing_enable) + + def test_trace_set_thread_info_disabled(self): + trace_set_thread_info("test_label") + # Should not register anything + + +@unittest.skipUnless(_has_otel, "opentelemetry not installed") +class TestTraceReqContextEnabled(unittest.TestCase): + def setUp(self): + + self.orig_initialized = mod.opentelemetry_initialized + self.orig_tracer = mod.tracer + self.orig_threads = mod.threads_info.copy() + self.orig_level = mod.global_trace_level + + self.provider = TracerProvider() + otel_trace.set_tracer_provider(self.provider) + mod.opentelemetry_initialized = True + mod.tracer = otel_trace.get_tracer("test") + mod.global_trace_level = 3 + + def tearDown(self): + mod.opentelemetry_initialized = self.orig_initialized + mod.tracer = self.orig_tracer + mod.threads_info.clear() + mod.threads_info.update(self.orig_threads) + mod.global_trace_level = self.orig_level + + def test_trace_set_thread_info(self): + trace_set_thread_info("scheduler", tp_rank=0, dp_rank=0) + + pid = threading.get_native_id() + self.assertIn(pid, mod.threads_info) + self.assertEqual(mod.threads_info[pid].thread_label, "scheduler") + + # Second call for same thread is a no-op + trace_set_thread_info("different_label") + self.assertEqual(mod.threads_info[pid].thread_label, "scheduler") + + def test_full_lifecycle(self): + """Start → slice_start → slice_end → finish.""" + ctx = TraceReqContext(rid="req-1", role="unified", module_name="test") + self.assertTrue(ctx.tracing_enable) + + ctx.trace_req_start(ts=1000) + self.assertEqual(ctx.start_time_ns, 1000) + self.assertIsNotNone(ctx.root_span) + self.assertIsNotNone(ctx.thread_context) + + ctx.trace_slice_start("prefill", level=1, ts=2000) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 1) + + ctx.trace_slice_end("prefill", level=1, ts=3000) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 0) + self.assertIsNotNone(ctx.last_span_context) + + ctx.trace_req_finish(ts=4000, attrs={"tokens": 42}) + self.assertIsNone(ctx.root_span) + + def test_trace_req_start_with_bootstrap_room(self): + ctx = TraceReqContext(rid="req-1", bootstrap_room=0xFF, role="prefill") + ctx.trace_req_start(ts=1000) + self.assertIsNotNone(ctx.root_span) + ctx.trace_req_finish(ts=2000) + + def test_trace_req_finish_without_start(self): + """finish without start is a no-op.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.root_span = None + ctx.trace_req_finish(ts=2000) + + def test_trace_slice_combined(self): + """trace_slice() creates and ends a span in one call.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + s = TraceSliceContext( + "decode", + 2000, + end_time_ns=3000, + level=1, + attrs={"key": "val"}, + events=[TraceEvent("evt", 2500, {"e": 1})], + ) + ctx.trace_slice(s) + self.assertIsNotNone(ctx.last_span_context) + ctx.trace_req_finish(ts=4000) + + def test_trace_slice_with_events_cache(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + # Add events to cache + ctx.trace_event("schedule", level=1, ts=1500, attrs={"bid": "x"}) + self.assertEqual(len(ctx.events_cache), 1) + + # trace_slice_start + trace_slice_end flushes matching events + ctx.trace_slice_start("prefill", level=1, ts=1200) + ctx.trace_slice_end("prefill", level=1, ts=2000) + self.assertEqual(len(ctx.events_cache), 0) + + ctx.trace_req_finish(ts=3000) + + def test_trace_slice_combined_with_events_cache(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + ctx.trace_event("evt", level=1, ts=1500) + s = TraceSliceContext("decode", 1200, end_time_ns=2000, level=1) + ctx.trace_slice(s) + self.assertEqual(len(ctx.events_cache), 0) + ctx.trace_req_finish(ts=3000) + + def test_trace_event_no_attrs(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_event("evt", level=1, ts=1500, attrs=None) + self.assertEqual(ctx.events_cache[0].attrs, {}) + ctx.trace_req_finish(ts=2000) + + def test_trace_slice_end_empty_stack(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + # End without start → warning, no crash + ctx.trace_slice_end("missing", level=1, ts=2000) + ctx.trace_req_finish(ts=3000) + + def test_trace_slice_end_name_mismatch(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("prefill", level=1, ts=1500) + # Mismatched name → warning, slice popped + ctx.trace_slice_end("wrong_name", level=1, ts=2000) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 0) + ctx.trace_req_finish(ts=3000) + + def test_trace_slice_end_with_attrs_and_thread_finish(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("dispatch", level=2, ts=1500) + ctx.trace_slice_end( + "dispatch", + level=2, + ts=2000, + attrs={"key": "val"}, + thread_finish_flag=True, + ) + # thread_finish_flag triggers abort → thread_context is None + self.assertIsNone(ctx.thread_context) + + def test_trace_slice_combined_with_thread_finish(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + s = TraceSliceContext("dispatch", 1500, end_time_ns=2000, level=2) + ctx.trace_slice(s, thread_finish_flag=True) + self.assertIsNone(ctx.thread_context) + + def test_nested_slices(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("outer", level=1, ts=1500) + ctx.trace_slice_start("inner", level=2, ts=1600) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 2) + ctx.trace_slice_end("inner", level=2, ts=1800) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 1) + ctx.trace_slice_end("outer", level=1, ts=2000) + ctx.trace_req_finish(ts=3000) + + def test_nested_slice_with_last_span_context(self): + """trace_slice uses last_span_context when slice stack is empty.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + # First slice sets last_span_context + ctx.trace_slice_start("s1", level=1, ts=1500) + ctx.trace_slice_end("s1", level=1, ts=2000) + self.assertIsNotNone(ctx.last_span_context) + + # Second slice uses last_span_context as link + ctx.trace_slice_start("s2", level=1, ts=2500) + ctx.trace_slice_end("s2", level=1, ts=3000) + + # trace_slice also uses last_span_context + s = TraceSliceContext("s3", 3500, end_time_ns=4000, level=1) + ctx.trace_slice(s) + + ctx.trace_req_finish(ts=5000) + + def test_trace_set_root_attrs(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_set_root_attrs({"model": "llama"}) + ctx.trace_req_finish(ts=2000) + + def test_trace_set_root_attrs_no_span(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.root_span = None + ctx.trace_set_root_attrs({"model": "llama"}) # no crash + + def test_trace_set_thread_attrs(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_set_thread_attrs({"batch_size": 32}) + ctx.trace_req_finish(ts=2000) + + def test_abort_with_unclosed_slices(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("s1", level=1, ts=1500) + ctx.trace_slice_start("s2", level=2, ts=1600) + ctx.abort(ts=2000) + self.assertIsNone(ctx.thread_context) + + def test_abort_with_events_cache(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_event("evt", level=1, ts=1500) + ctx.abort(ts=2000) + self.assertEqual(len(ctx.events_cache), 0) + + def test_abort_with_abort_info_dict(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.abort(ts=2000, abort_info={"reason": "cancelled"}) + self.assertIsNone(ctx.thread_context) + + def test_abort_with_base_finish_reason(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + abort_obj = _sb.BaseFinishReason() + ctx.abort(ts=2000, abort_info=abort_obj) + self.assertIsNone(ctx.thread_context) + + def test_check_fast_return_by_level(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_level = 1 # instance-level, set at init from global + # Level 2 > trace_level 1 → fast return + ctx.trace_slice_start("s", level=2, ts=1500) + self.assertEqual(len(ctx.thread_context.cur_slice_stack), 0) + ctx.trace_level = 3 + ctx.trace_req_finish(ts=2000) + + def test_rebuild_thread_context(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + old_tc = ctx.thread_context + ctx.rebuild_thread_context(ts=1500) + self.assertIsNot(ctx.thread_context, old_tc) + ctx.trace_req_finish(ts=2000) + + def test_getstate_enabled(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + state = ctx.__getstate__() + self.assertTrue(state["tracing_enable"]) + self.assertEqual(state["rid"], "req-1") + self.assertIn("root_span_context", state) + ctx.trace_req_finish(ts=2000) + + def test_getstate_no_root_context(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.root_span_context = None + state = ctx.__getstate__() + self.assertFalse(state["tracing_enable"]) + ctx.root_span_context = True # prevent __del__ issues + ctx.trace_req_finish(ts=2000) + + def test_getstate_with_slice_stack(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("s1", level=1, ts=1500) + state = ctx.__getstate__() + self.assertIn("last_span_context", state) + ctx.trace_req_finish(ts=2000) + + def test_setstate_enabled(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + state = ctx.__getstate__() + ctx.trace_req_finish(ts=2000) + + ctx2 = TraceReqContext.__new__(TraceReqContext) + ctx2.__setstate__(state) + self.assertTrue(ctx2.tracing_enable) + self.assertTrue(ctx2.is_copy) + self.assertIsNotNone(ctx2.root_span_context) + + def test_thread_context_with_tp_rank(self): + """Covers tp_rank branch in __create_thread_context.""" + + pid = threading.get_native_id() + mod.threads_info[pid] = TraceThreadInfo( + "host", pid, "sched", tp_rank=0, dp_rank=0 + ) + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + self.assertIsNotNone(ctx.thread_context) + ctx.trace_req_finish(ts=2000) + + def test_setstate_with_last_span_context(self): + """Covers __setstate__ path where last_span_context is truthy.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + ctx.trace_slice_start("s1", level=1, ts=1500) + ctx.trace_slice_end("s1", level=1, ts=2000) + state = ctx.__getstate__() + ctx.trace_req_finish(ts=3000) + + self.assertIsNotNone(state.get("last_span_context")) + ctx2 = TraceReqContext.__new__(TraceReqContext) + ctx2.__setstate__(state) + self.assertIsNotNone(ctx2.last_span_context) + + def test_events_cache_partial_match(self): + """Events outside the slice time range stay in cache.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + ctx.trace_event("early", level=1, ts=500) + ctx.trace_event("inside", level=1, ts=1500) + ctx.trace_event("late", level=1, ts=5000) + + ctx.trace_slice_start("s", level=1, ts=1200) + ctx.trace_slice_end("s", level=1, ts=2000) + # "early" (500 < 1200) and "late" (5000 >= 2000) stay in cache + self.assertEqual(len(ctx.events_cache), 2) + ctx.trace_req_finish(ts=6000) + + def test_trace_slice_combined_events_partial_match(self): + """Events outside slice range stay in cache for trace_slice method.""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + ctx.trace_event("early", level=1, ts=500) + ctx.trace_event("inside", level=1, ts=1500) + + s = TraceSliceContext("s", 1200, end_time_ns=2000, level=1) + ctx.trace_slice(s) + self.assertEqual(len(ctx.events_cache), 1) # "early" stays + ctx.trace_req_finish(ts=3000) + + def test_trace_slice_nested_parent(self): + """trace_slice with parent from slice stack (not thread_span).""" + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + + ctx.trace_slice_start("outer", level=1, ts=1500) + s = TraceSliceContext("inner", 1600, end_time_ns=1800, level=2) + ctx.trace_slice(s) + ctx.trace_slice_end("outer", level=1, ts=2000) + ctx.trace_req_finish(ts=3000) + + def test_del_triggers_abort(self): + ctx = TraceReqContext(rid="req-1") + ctx.trace_req_start(ts=1000) + # __del__ calls abort + ctx.__del__() + self.assertIsNone(ctx.thread_context) + + +if __name__ == "__main__": + unittest.main()