[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:
co-authored by
Mu Huai
Shangming Cai
parent
73b53e7a87
commit
e67810bea7
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user