diff --git a/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py b/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py index 8288261d3..6dc5c987a 100644 --- a/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py +++ b/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py @@ -32,32 +32,17 @@ def init_diffusion_tracing(server_args, thread_label: str): if not server_args.enable_trace: return - from types import SimpleNamespace - - from sglang.srt import server_args as srt_server_args_module from sglang.srt.observability.trace import ( process_tracing_init, trace_set_thread_info, ) - from sglang.srt.server_args import set_global_server_args_for_scheduler - # srt owns TraceReqContext and filters spans through its global trace_modules - try: - srt_server_args = srt_server_args_module.get_global_server_args() - except ValueError: - srt_server_args = SimpleNamespace(trace_modules=DIFFUSION_TRACE_MODULE) - set_global_server_args_for_scheduler(srt_server_args) - - trace_modules = [ - module.strip() - for module in getattr(srt_server_args, "trace_modules", "").split(",") - if module.strip() - ] - if DIFFUSION_TRACE_MODULE not in trace_modules: - trace_modules.append(DIFFUSION_TRACE_MODULE) - srt_server_args.trace_modules = ",".join(trace_modules) - - process_tracing_init(server_args.otlp_traces_endpoint, "sglang-diffusion") + # srt owns TraceReqContext and filters spans through its trace_modules list + process_tracing_init( + server_args.otlp_traces_endpoint, + "sglang-diffusion", + trace_modules=DIFFUSION_TRACE_MODULE, + ) trace_set_thread_info(thread_label) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index cf1bf6a57..d17e93d15 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -257,7 +257,11 @@ class Engine(EngineScoreMixin, EngineBase): # Enable tracing if server_args.enable_trace: - process_tracing_init(server_args.otlp_traces_endpoint, "sglang") + process_tracing_init( + server_args.otlp_traces_endpoint, + "sglang", + trace_modules=server_args.trace_modules, + ) thread_label = "Tokenizer" if server_args.disaggregation_mode == "prefill": thread_label = "Prefill Tokenizer" diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index c51dbcd2f..fdacd0ad5 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -306,7 +306,11 @@ async def lifespan(fast_api_app: FastAPI): # Init tracing if server_args.enable_trace: - process_tracing_init(server_args.otlp_traces_endpoint, "sglang") + process_tracing_init( + server_args.otlp_traces_endpoint, + "sglang", + trace_modules=server_args.trace_modules, + ) if server_args.disaggregation_mode == "prefill": thread_label = "Prefill" + thread_label elif server_args.disaggregation_mode == "decode": diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index e84a593bf..5b5bc4ed6 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -667,7 +667,11 @@ def run_data_parallel_controller_process( configure_logger(server_args) if server_args.enable_trace: - process_tracing_init(server_args.otlp_traces_endpoint, "sglang") + process_tracing_init( + server_args.otlp_traces_endpoint, + "sglang", + trace_modules=server_args.trace_modules, + ) thread_label = "DP Controller" if server_args.disaggregation_mode == "prefill": thread_label = "Prefill DP Controller" diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index ac316b716..ea515de96 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3977,7 +3977,11 @@ def run_scheduler_process( # Set up tracing if server_args.enable_trace: - process_tracing_init(server_args.otlp_traces_endpoint, "sglang") + process_tracing_init( + server_args.otlp_traces_endpoint, + "sglang", + trace_modules=server_args.trace_modules, + ) thread_label = "Scheduler" if server_args.disaggregation_mode == "prefill": thread_label = "Prefill Scheduler" diff --git a/python/sglang/srt/observability/trace.py b/python/sglang/srt/observability/trace.py index c959f7a8d..7c38b6926 100644 --- a/python/sglang/srt/observability/trace.py +++ b/python/sglang/srt/observability/trace.py @@ -22,14 +22,10 @@ import threading import time import uuid from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional +from typing import Any, Dict, List, Mapping, Optional -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import get_int_env_var -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - logger = logging.getLogger(__name__) opentelemetry_imported = False opentelemetry_initialized = False @@ -38,6 +34,9 @@ tracer: Optional[trace.Tracer] = None global_trace_level = get_int_env_var("SGLANG_TRACE_LEVEL", 3) +# Modules allowed to emit spans (from --trace-modules); None means no filtering. +global_trace_modules: Optional[List[str]] = None + TRACE_HEADERS = ["traceparent", "tracestate"] try: @@ -162,10 +161,19 @@ def _get_host_id() -> str: # Should be called by each tracked process. -def process_tracing_init(otlp_endpoint, server_name): +def process_tracing_init( + otlp_endpoint, server_name, trace_modules: Optional[str] = None +): global opentelemetry_initialized global get_cur_time_ns global tracer + global global_trace_modules + + if trace_modules is not None: + global_trace_modules = [ + module.strip() for module in trace_modules.split(",") if module.strip() + ] + if not opentelemetry_imported: opentelemetry_initialized = False raise RuntimeError( @@ -263,8 +271,13 @@ class TraceReqContext: self.trace_level = global_trace_level self.tracing_enable: bool = opentelemetry_initialized and self.trace_level > 0 - server_args: ServerArgs = get_global_server_args() - if module_name not in server_args.trace_modules.split(","): + # Filter by --trace-modules only for explicitly named modules; contexts + # created with the default empty module_name are always traced. + if ( + module_name + and global_trace_modules is not None + and module_name not in global_trace_modules + ): self.tracing_enable = False if not self.tracing_enable: diff --git a/test/registered/unit/observability/test_trace.py b/test/registered/unit/observability/test_trace.py index e26d9bab0..7667bce92 100644 --- a/test/registered/unit/observability/test_trace.py +++ b/test/registered/unit/observability/test_trace.py @@ -8,7 +8,7 @@ register_cpu_ci(est_time=6, suite="base-a-test-cpu") import threading import unittest -from unittest.mock import MagicMock, patch +from unittest.mock import patch import sglang.srt.observability.trace as mod from sglang.srt.observability.trace import ( @@ -195,25 +195,12 @@ 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): @@ -269,14 +256,7 @@ class TestTraceReqContextEnabled(unittest.TestCase): 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() @@ -294,6 +274,23 @@ class TestTraceReqContextEnabled(unittest.TestCase): trace_set_thread_info("different_label") self.assertEqual(mod.threads_info[pid].thread_label, "scheduler") + def test_module_filtering(self): + """global_trace_modules gates only explicitly named modules.""" + orig_modules = mod.global_trace_modules + mod.global_trace_modules = ["request"] + try: + # Default empty module_name is never filtered + ctx = TraceReqContext(rid="req-1") + self.assertTrue(ctx.tracing_enable) + # Listed module is traced + ctx = TraceReqContext(rid="req-1", module_name="request") + self.assertTrue(ctx.tracing_enable) + # Unlisted module is filtered out + ctx = TraceReqContext(rid="req-1", module_name="mooncake") + self.assertFalse(ctx.tracing_enable) + finally: + mod.global_trace_modules = orig_modules + def test_full_lifecycle(self): """Start → slice_start → slice_end → finish.""" ctx = TraceReqContext(rid="req-1", role="unified")