Fix trace_modules gate disabling default trace contexts (#27173)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user