[SGLang Tracing] Add pd disaggregation mooncake backend tracing (#23755)

Co-authored-by: Mu Huai <tianbowen.tbw@antgroup.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Feng Su
2026-06-03 16:43:29 +08:00
committed by GitHub
co-authored by Mu Huai Shangming Cai
parent 73b53e7a87
commit e67810bea7
9 changed files with 270 additions and 14 deletions
@@ -15,6 +15,7 @@ from urllib.parse import urlparse
import requests
from sglang.srt.observability.mooncake_trace import MooncakeRequestStage
from sglang.srt.observability.req_time_stats import RequestStage
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
@@ -78,6 +79,8 @@ class TestTraceDisaggregation(CustomTestCase):
"--enable-trace",
"--otlp-traces-endpoint",
"localhost:4317",
"--trace-modules",
"request,mooncake",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
@@ -186,9 +189,7 @@ class TestTraceDisaggregation(CustomTestCase):
def test_disaggregation_transfer_spans(self):
"""Test that disaggregation produces PREFILL_TRANSFER_KV_CACHE and DECODE_TRANSFERRED spans."""
# Set trace level
response = requests.get(f"{self.prefill_url}/set_trace_level?level=1")
self.assertEqual(response.status_code, 200)
response = requests.get(f"{self.decode_url}/set_trace_level?level=1")
response = requests.get(f"{self.prefill_url}/set_trace_level?level=2")
self.assertEqual(response.status_code, 200)
self.collector.clear()
@@ -221,13 +222,14 @@ class TestTraceDisaggregation(CustomTestCase):
# Check for transfer-related spans
self.assertTrue(
self.collector.has_any_span(
self.collector.has_all_spans(
[
RequestStage.PREFILL_TRANSFER_KV_CACHE.stage_name,
RequestStage.DECODE_TRANSFERRED.stage_name,
MooncakeRequestStage.MOONCAKE_WORKER_SEND.stage_name,
]
),
f"Expected disaggregation transfer spans, got {sorted(span_names)}",
f"Expected all disaggregation transfer spans, got {sorted(span_names)}",
)
@@ -8,7 +8,7 @@ register_cpu_ci(est_time=6, suite="base-a-test-cpu")
import threading
import unittest
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import sglang.srt.observability.trace as mod
from sglang.srt.observability.trace import (
@@ -38,7 +38,7 @@ except ImportError:
_has_otel = False
# Access the private module-level function (avoid name mangling inside classes).
_get_host_id = getattr(mod, "__get_host_id")
_get_host_id = getattr(mod, "_get_host_id")
class TestTraceFunctions(unittest.TestCase):
@@ -195,12 +195,25 @@ class TestProcessTracingInit(unittest.TestCase):
mod.opentelemetry_imported = orig
def _mock_get_global_server_args():
"""Return a mock ServerArgs for tests that create TraceReqContext."""
mock = MagicMock()
mock.trace_modules = ""
return mock
class TestTraceReqContextDisabled(unittest.TestCase):
def setUp(self):
self.orig = mod.opentelemetry_initialized
mod.opentelemetry_initialized = False
self._sa_patcher = patch(
"sglang.srt.observability.trace.get_global_server_args",
side_effect=_mock_get_global_server_args,
)
self._sa_patcher.start()
def tearDown(self):
self._sa_patcher.stop()
mod.opentelemetry_initialized = self.orig
def test_init_disabled(self):
@@ -246,13 +259,24 @@ class TestTraceReqContextEnabled(unittest.TestCase):
self.orig_threads = mod.threads_info.copy()
self.orig_level = mod.global_trace_level
# Reset OTel global TracerProvider so set_tracer_provider works each test
otel_trace._TRACER_PROVIDER_SET_ONCE._done = False
otel_trace._TRACER_PROVIDER = None
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
self._sa_patcher = patch(
"sglang.srt.observability.trace.get_global_server_args",
side_effect=_mock_get_global_server_args,
)
self._sa_patcher.start()
def tearDown(self):
self._sa_patcher.stop()
mod.opentelemetry_initialized = self.orig_initialized
mod.tracer = self.orig_tracer
mod.threads_info.clear()
@@ -272,7 +296,7 @@ class TestTraceReqContextEnabled(unittest.TestCase):
def test_full_lifecycle(self):
"""Start → slice_start → slice_end → finish."""
ctx = TraceReqContext(rid="req-1", role="unified", module_name="test")
ctx = TraceReqContext(rid="req-1", role="unified")
self.assertTrue(ctx.tracing_enable)
ctx.trace_req_start(ts=1000)