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:
|
if not server_args.enable_trace:
|
||||||
return
|
return
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
from sglang.srt import server_args as srt_server_args_module
|
|
||||||
from sglang.srt.observability.trace import (
|
from sglang.srt.observability.trace import (
|
||||||
process_tracing_init,
|
process_tracing_init,
|
||||||
trace_set_thread_info,
|
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
|
# srt owns TraceReqContext and filters spans through its trace_modules list
|
||||||
try:
|
process_tracing_init(
|
||||||
srt_server_args = srt_server_args_module.get_global_server_args()
|
server_args.otlp_traces_endpoint,
|
||||||
except ValueError:
|
"sglang-diffusion",
|
||||||
srt_server_args = SimpleNamespace(trace_modules=DIFFUSION_TRACE_MODULE)
|
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")
|
|
||||||
trace_set_thread_info(thread_label)
|
trace_set_thread_info(thread_label)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -257,7 +257,11 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
|
|
||||||
# Enable tracing
|
# Enable tracing
|
||||||
if server_args.enable_trace:
|
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"
|
thread_label = "Tokenizer"
|
||||||
if server_args.disaggregation_mode == "prefill":
|
if server_args.disaggregation_mode == "prefill":
|
||||||
thread_label = "Prefill Tokenizer"
|
thread_label = "Prefill Tokenizer"
|
||||||
|
|||||||
@@ -306,7 +306,11 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
|
|
||||||
# Init tracing
|
# Init tracing
|
||||||
if server_args.enable_trace:
|
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":
|
if server_args.disaggregation_mode == "prefill":
|
||||||
thread_label = "Prefill" + thread_label
|
thread_label = "Prefill" + thread_label
|
||||||
elif server_args.disaggregation_mode == "decode":
|
elif server_args.disaggregation_mode == "decode":
|
||||||
|
|||||||
@@ -667,7 +667,11 @@ def run_data_parallel_controller_process(
|
|||||||
|
|
||||||
configure_logger(server_args)
|
configure_logger(server_args)
|
||||||
if server_args.enable_trace:
|
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"
|
thread_label = "DP Controller"
|
||||||
if server_args.disaggregation_mode == "prefill":
|
if server_args.disaggregation_mode == "prefill":
|
||||||
thread_label = "Prefill DP Controller"
|
thread_label = "Prefill DP Controller"
|
||||||
|
|||||||
@@ -3977,7 +3977,11 @@ def run_scheduler_process(
|
|||||||
|
|
||||||
# Set up tracing
|
# Set up tracing
|
||||||
if server_args.enable_trace:
|
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"
|
thread_label = "Scheduler"
|
||||||
if server_args.disaggregation_mode == "prefill":
|
if server_args.disaggregation_mode == "prefill":
|
||||||
thread_label = "Prefill Scheduler"
|
thread_label = "Prefill Scheduler"
|
||||||
|
|||||||
@@ -22,14 +22,10 @@ import threading
|
|||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass
|
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
|
from sglang.srt.utils import get_int_env_var
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
opentelemetry_imported = False
|
opentelemetry_imported = False
|
||||||
opentelemetry_initialized = False
|
opentelemetry_initialized = False
|
||||||
@@ -38,6 +34,9 @@ tracer: Optional[trace.Tracer] = None
|
|||||||
|
|
||||||
global_trace_level = get_int_env_var("SGLANG_TRACE_LEVEL", 3)
|
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"]
|
TRACE_HEADERS = ["traceparent", "tracestate"]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -162,10 +161,19 @@ def _get_host_id() -> str:
|
|||||||
|
|
||||||
|
|
||||||
# Should be called by each tracked process.
|
# 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 opentelemetry_initialized
|
||||||
global get_cur_time_ns
|
global get_cur_time_ns
|
||||||
global tracer
|
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:
|
if not opentelemetry_imported:
|
||||||
opentelemetry_initialized = False
|
opentelemetry_initialized = False
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -263,8 +271,13 @@ class TraceReqContext:
|
|||||||
self.trace_level = global_trace_level
|
self.trace_level = global_trace_level
|
||||||
self.tracing_enable: bool = opentelemetry_initialized and self.trace_level > 0
|
self.tracing_enable: bool = opentelemetry_initialized and self.trace_level > 0
|
||||||
|
|
||||||
server_args: ServerArgs = get_global_server_args()
|
# Filter by --trace-modules only for explicitly named modules; contexts
|
||||||
if module_name not in server_args.trace_modules.split(","):
|
# 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
|
self.tracing_enable = False
|
||||||
|
|
||||||
if not self.tracing_enable:
|
if not self.tracing_enable:
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
|||||||
|
|
||||||
import threading
|
import threading
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import sglang.srt.observability.trace as mod
|
import sglang.srt.observability.trace as mod
|
||||||
from sglang.srt.observability.trace import (
|
from sglang.srt.observability.trace import (
|
||||||
@@ -195,25 +195,12 @@ class TestProcessTracingInit(unittest.TestCase):
|
|||||||
mod.opentelemetry_imported = orig
|
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):
|
class TestTraceReqContextDisabled(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.orig = mod.opentelemetry_initialized
|
self.orig = mod.opentelemetry_initialized
|
||||||
mod.opentelemetry_initialized = False
|
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):
|
def tearDown(self):
|
||||||
self._sa_patcher.stop()
|
|
||||||
mod.opentelemetry_initialized = self.orig
|
mod.opentelemetry_initialized = self.orig
|
||||||
|
|
||||||
def test_init_disabled(self):
|
def test_init_disabled(self):
|
||||||
@@ -269,14 +256,7 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
|||||||
mod.tracer = otel_trace.get_tracer("test")
|
mod.tracer = otel_trace.get_tracer("test")
|
||||||
mod.global_trace_level = 3
|
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):
|
def tearDown(self):
|
||||||
self._sa_patcher.stop()
|
|
||||||
mod.opentelemetry_initialized = self.orig_initialized
|
mod.opentelemetry_initialized = self.orig_initialized
|
||||||
mod.tracer = self.orig_tracer
|
mod.tracer = self.orig_tracer
|
||||||
mod.threads_info.clear()
|
mod.threads_info.clear()
|
||||||
@@ -294,6 +274,23 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
|||||||
trace_set_thread_info("different_label")
|
trace_set_thread_info("different_label")
|
||||||
self.assertEqual(mod.threads_info[pid].thread_label, "scheduler")
|
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):
|
def test_full_lifecycle(self):
|
||||||
"""Start → slice_start → slice_end → finish."""
|
"""Start → slice_start → slice_end → finish."""
|
||||||
ctx = TraceReqContext(rid="req-1", role="unified")
|
ctx = TraceReqContext(rid="req-1", role="unified")
|
||||||
|
|||||||
Reference in New Issue
Block a user