ci: unit test for srt/observability module (#21002)

This commit is contained in:
Jiabin Wang
2026-03-24 00:30:02 +08:00
committed by GitHub
parent 4dbe42527e
commit 83b9d74424
7 changed files with 2088 additions and 0 deletions
@@ -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()
@@ -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()
@@ -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()
@@ -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()
@@ -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()
@@ -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()
@@ -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()