diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx
index 6cf6d6343..f63cf7b52 100644
--- a/docs_new/docs/advanced_features/server_arguments.mdx
+++ b/docs_new/docs/advanced_features/server_arguments.mdx
@@ -814,6 +814,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s
False |
bool flag (set to enable) |
+
+ --trace-modules |
+ Select the components to trace. Available options are 'request' and 'mooncake'. Format: <module1 name>,<module2 name>,...... |
+ request |
+ Type: str |
+
| `--otlp-traces-endpoint` |
Config opentelemetry collector endpoint if --enable-trace is set. format: <ip>:<port> |
diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py
index fadce5383..e4571f2d9 100644
--- a/python/sglang/srt/disaggregation/common/utils.py
+++ b/python/sglang/srt/disaggregation/common/utils.py
@@ -3,11 +3,16 @@ import dataclasses
import struct
import threading
from collections import deque
-from typing import List, Optional, Tuple
+from typing import List, Optional, Tuple, Union
import numpy as np
import numpy.typing as npt
+from sglang.srt.observability.trace import (
+ TraceNullContext,
+ TraceReqContext,
+)
+
@dataclasses.dataclass
class TransferKVChunk:
@@ -20,6 +25,9 @@ class TransferKVChunk:
prefill_aux_index: Optional[int]
state_indices: Optional[List]
chunk_id: Optional[int] = None
+ trace_ctx: Union[TraceReqContext, TraceNullContext] = dataclasses.field(
+ default_factory=TraceNullContext
+ )
def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py
index c70d6fa51..3d4406452 100644
--- a/python/sglang/srt/disaggregation/mooncake/conn.py
+++ b/python/sglang/srt/disaggregation/mooncake/conn.py
@@ -8,7 +8,7 @@ import struct
import threading
import time
from collections import defaultdict
-from typing import List, Optional, Tuple
+from typing import List, Optional, Tuple, Union
import numpy as np
import numpy.typing as npt
@@ -41,6 +41,16 @@ from sglang.srt.disaggregation.mooncake.utils import (
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
from sglang.srt.environ import envs
+from sglang.srt.observability.mooncake_trace import (
+ MooncakeRequestStage,
+ mooncake_trace_func,
+ mooncake_trace_slice,
+)
+from sglang.srt.observability.trace import (
+ TraceNullContext,
+ TraceReqContext,
+ trace_set_thread_info,
+)
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import NetworkAddress
@@ -163,6 +173,7 @@ class MooncakeKVManager(CommonKVManager):
self.init_engine()
self.register_buffer_to_engine()
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
+ self.enable_trace = server_args.enable_trace
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self.start_prefill_thread()
self.session_failures = defaultdict(int)
@@ -208,6 +219,7 @@ class MooncakeKVManager(CommonKVManager):
if self.enable_staging and self._staging_ctx.buffers
else None
),
+ i,
),
daemon=True,
).start()
@@ -1136,12 +1148,26 @@ class MooncakeKVManager(CommonKVManager):
queue: FastQueue,
executor: concurrent.futures.ThreadPoolExecutor,
staging_buffer=None,
+ worker_index=0,
):
staging_strategy = None
+ if self.enable_trace:
+ trace_set_thread_info(
+ f"mooncake transfer worker {worker_index}",
+ tp_rank=self.attn_tp_rank,
+ dp_rank=self.attn_dp_rank,
+ )
while True:
try:
kv_chunk: TransferKVChunk = queue.get()
+ if self.enable_trace:
+ kv_chunk.trace_ctx.rebuild_thread_context()
+ kv_chunk.trace_ctx.trace_slice_start(
+ MooncakeRequestStage.MOONCAKE_WORKER_SEND.stage_name,
+ MooncakeRequestStage.MOONCAKE_WORKER_SEND.level,
+ )
+
if (
self.enable_staging
and staging_strategy is None
@@ -1165,6 +1191,7 @@ class MooncakeKVManager(CommonKVManager):
# the chunk is re-enqueued and we break out of the req loop to retry later.
staging_deferred = False
for req in reqs_to_be_processed:
+ start_ts = time.perf_counter()
if not req.is_dummy:
# Early exit if the request has failed
with self.session_lock:
@@ -1305,6 +1332,20 @@ class MooncakeKVManager(CommonKVManager):
if kv_chunk.is_last_chunk and req.room in self.request_status:
self.update_status(req.room, KVPoll.Success)
+ if self.enable_trace:
+ mooncake_trace_slice(
+ kv_chunk.trace_ctx,
+ MooncakeRequestStage.MOONCAKE_WORKER_SEND_SESSION,
+ start_ts,
+ )
+
+ if self.enable_trace:
+ kv_chunk.trace_ctx.trace_slice_end(
+ MooncakeRequestStage.MOONCAKE_WORKER_SEND.stage_name,
+ MooncakeRequestStage.MOONCAKE_WORKER_SEND.level,
+ thread_finish_flag=True,
+ )
+
if staging_deferred:
continue
@@ -1455,6 +1496,7 @@ class MooncakeKVManager(CommonKVManager):
is_last_chunk: bool,
aux_index: Optional[int] = None,
state_indices: Optional[List] = None,
+ trace_ctx: Optional[Union[TraceReqContext, TraceNullContext]] = None,
):
assert self.disaggregation_mode == DisaggregationMode.PREFILL
assert not is_last_chunk or (is_last_chunk and aux_index is not None)
@@ -1481,6 +1523,9 @@ class MooncakeKVManager(CommonKVManager):
session_port_sum = sum(int(session.rsplit(":", 1)[1]) for session in dst_infos)
shard_idx = session_port_sum % len(self.transfer_queues)
+ if trace_ctx is None:
+ trace_ctx = TraceNullContext()
+
self.transfer_queues[shard_idx].put(
TransferKVChunk(
room=bootstrap_room,
@@ -1489,6 +1534,7 @@ class MooncakeKVManager(CommonKVManager):
is_last_chunk=is_last_chunk,
prefill_aux_index=aux_index,
state_indices=state_indices,
+ trace_ctx=trace_ctx,
)
)
@@ -1584,7 +1630,9 @@ class MooncakeKVSender(CommonKVSender):
super().__init__(mgr, bootstrap_addr, bootstrap_room, dest_tp_ranks, pp_rank)
self.conclude_state = None
self.init_time = time.time()
+ self._init_trace_ctx()
+ @mooncake_trace_func(MooncakeRequestStage.MOONCAKE_SEND)
def send(
self,
kv_indices: npt.NDArray[np.int32],
@@ -1602,6 +1650,7 @@ class MooncakeKVSender(CommonKVSender):
kv_indices,
index_slice,
False,
+ trace_ctx=self.trace_ctx.copy_for_thread(),
)
else:
self.kv_mgr.add_transfer_request(
@@ -1611,6 +1660,7 @@ class MooncakeKVSender(CommonKVSender):
True,
aux_index=self.aux_index,
state_indices=state_indices,
+ trace_ctx=self.trace_ctx.copy_for_thread(),
)
self._record_transfer_indices(kv_indices, state_indices)
@@ -1619,6 +1669,7 @@ class MooncakeKVSender(CommonKVSender):
status = self.kv_mgr.check_status(self.bootstrap_room)
if status in (KVPoll.Success, KVPoll.Failed):
self.conclude_state = status
+ self.trace_ctx.trace_req_finish()
elif status == KVPoll.Bootstrapping:
timeout_result = self._check_bootstrap_timeout()
if timeout_result is not None:
@@ -1641,6 +1692,26 @@ class MooncakeKVSender(CommonKVSender):
)
raise KVTransferError(self.bootstrap_room, failure_reason)
+ def _init_trace_ctx(self):
+ if self.kv_mgr.enable_trace:
+ self.trace_ctx = TraceReqContext(
+ rid=str(hex(self.bootstrap_room)),
+ bootstrap_room=self.bootstrap_room,
+ role="Sender",
+ module_name="mooncake",
+ )
+ if not self.trace_ctx.tracing_enable:
+ self.trace_ctx = TraceNullContext()
+ else:
+ self.trace_ctx = TraceNullContext()
+
+ self.trace_ctx.trace_req_start()
+
+ def abort(self):
+ super().abort()
+ self.trace_ctx.abort(abort_info={"reason": "Aborted"})
+ self.trace_ctx.trace_req_finish()
+
class MooncakeKVReceiver(CommonKVReceiver):
def __init__(
diff --git a/python/sglang/srt/observability/mooncake_trace.py b/python/sglang/srt/observability/mooncake_trace.py
new file mode 100644
index 000000000..fdb110ce5
--- /dev/null
+++ b/python/sglang/srt/observability/mooncake_trace.py
@@ -0,0 +1,68 @@
+import time
+from typing import Union
+
+from sglang.srt.observability.req_time_stats import (
+ RequestStageConfig,
+ convert_time_to_realtime_ns,
+)
+from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
+
+
+class MooncakeRequestStage:
+ MOONCAKE_SEND = RequestStageConfig(
+ "mooncake_send",
+ level=1,
+ )
+ MOONCAKE_RECV = RequestStageConfig(
+ "mooncake_recv",
+ level=1,
+ )
+ MOONCAKE_WORKER_SEND = RequestStageConfig(
+ "mooncake_worker_send",
+ level=1,
+ )
+ MOONCAKE_WORKER_SEND_SESSION = RequestStageConfig(
+ "mooncake_worker_send_session",
+ level=2,
+ )
+ MOONCAKE_WORKER_RECV = RequestStageConfig(
+ "mooncake_worker_recv",
+ level=1,
+ )
+
+
+def mooncake_trace_slice(
+ trace_ctx: Union[TraceReqContext, TraceNullContext],
+ stage: RequestStageConfig,
+ start_ts: float,
+ thread_finish_flag=False,
+):
+ if trace_ctx is None:
+ return
+
+ if not trace_ctx.tracing_enable:
+ return
+
+ start_ts = convert_time_to_realtime_ns(start_ts)
+ trace_ctx.trace_slice_start(stage.stage_name, stage.level, start_ts)
+ trace_ctx.trace_slice_end(
+ stage.stage_name,
+ stage.level,
+ thread_finish_flag=thread_finish_flag,
+ )
+
+
+def mooncake_trace_func(stage: RequestStageConfig):
+ def decorator(func):
+ def wrapper(self, *args, **kwargs):
+ if self.trace_ctx is None:
+ return func(self, *args, **kwargs)
+ start_ts = convert_time_to_realtime_ns(time.perf_counter())
+ self.trace_ctx.trace_slice_start(stage.stage_name, stage.level, start_ts)
+ ret = func(self, *args, **kwargs)
+ self.trace_ctx.trace_slice_end(stage.stage_name, stage.level)
+ return ret
+
+ return wrapper
+
+ return decorator
diff --git a/python/sglang/srt/observability/trace.py b/python/sglang/srt/observability/trace.py
index a7cc4cfdd..c959f7a8d 100644
--- a/python/sglang/srt/observability/trace.py
+++ b/python/sglang/srt/observability/trace.py
@@ -22,10 +22,14 @@ import threading
import time
import uuid
from dataclasses import dataclass
-from typing import Any, Dict, List, Mapping, Optional
+from typing import TYPE_CHECKING, 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
@@ -137,7 +141,7 @@ if hasattr(time, "time_ns"):
get_cur_time_ns = lambda: int(time.time_ns())
-def __get_host_id() -> str:
+def _get_host_id() -> str:
"""
In distributed tracing systems, obtain a unique node identifier
and inject it into all subsequently generated spans
@@ -237,7 +241,7 @@ def trace_set_thread_info(
return
threads_info[pid] = TraceThreadInfo(
- host_id=__get_host_id(),
+ host_id=_get_host_id(),
pid=pid,
thread_label=thread_label,
tp_rank=tp_rank,
@@ -259,6 +263,10 @@ 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(","):
+ self.tracing_enable = False
+
if not self.tracing_enable:
return
@@ -386,6 +394,66 @@ class TraceReqContext:
)
self.events_cache = []
+ def copy_for_thread(self) -> "TraceReqContext":
+ """
+ Create a copy of this context for use in another thread.
+
+ The copy shares the same root_span_context but has its own thread_context.
+ This is useful for propagating trace context across threads (e.g., worker threads).
+
+ Usage:
+ # Sender (main thread)
+ trace_ctx_copy = trace_ctx.copy_for_thread()
+ queue.put(TransferKVChunk(..., trace_ctx=trace_ctx_copy))
+
+ # Receiver (worker thread)
+ kv_chunk = queue.get()
+ kv_chunk.trace_ctx.rebuild_thread_context()
+ """
+ # Fast path: not tracing
+ if not self.tracing_enable or not self.root_span_context:
+ return TraceNullContext()
+
+ # Extract prev_span_context from current thread state
+ prev_span_context = self.last_span_context
+ if self.thread_context and self.thread_context.cur_slice_stack:
+ cur_slice = self.thread_context.cur_slice_stack[0]
+ if cur_slice.span:
+ prev_span_context = cur_slice.span.get_span_context()
+
+ # Create new instance with shared state
+ copied = TraceReqContext.__new__(TraceReqContext)
+ copied.tracing_enable = self.tracing_enable
+ copied.rid = self.rid
+ copied.bootstrap_room = self.bootstrap_room
+ copied.start_time_ns = self.start_time_ns
+ copied.role = self.role
+ copied.trace_level = self.trace_level
+ copied.module_name = self.module_name
+ copied.is_copy = True # Mark as copy
+ copied.pid = self.pid
+
+ # thread_context is None, will be rebuilt via rebuild_thread_context()
+ copied.thread_context = None
+ copied.root_span = None
+
+ # Share root_span_context (already a context, no need to serialize)
+ copied.root_span_context = self.root_span_context
+
+ # Set prev_span_context for linking spans
+ if prev_span_context:
+ copied.last_span_context = trace.span.SpanContext(
+ trace_id=prev_span_context.trace_id,
+ span_id=prev_span_context.span_id,
+ is_remote=True,
+ )
+ else:
+ copied.last_span_context = None
+
+ copied.events_cache = []
+
+ return copied
+
def rebuild_thread_context(self, ts: Optional[int] = None):
if not self.tracing_enable:
return
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index a0d707fa3..5cb20a067 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -494,6 +494,7 @@ class ServerArgs:
forward_pass_metrics_worker_id: str = ""
forward_pass_metrics_ipc_name: Optional[str] = None
enable_trace: bool = False
+ trace_modules: str = "request"
otlp_traces_endpoint: str = "localhost:4317"
# RequestMetricsExporter configuration
@@ -5258,6 +5259,13 @@ class ServerArgs:
action="store_true",
help="Enable opentelemetry trace",
)
+ parser.add_argument(
+ "--trace-modules",
+ type=str,
+ default="request",
+ help="Select the components to trace. Available options are 'request' and 'mooncake'. Format: ,,...",
+ )
+
parser.add_argument(
"--otlp-traces-endpoint",
type=str,
diff --git a/scripts/convert_otel_2_perfetto.py b/scripts/convert_otel_2_perfetto.py
index 89534a38f..39ad3f78e 100644
--- a/scripts/convert_otel_2_perfetto.py
+++ b/scripts/convert_otel_2_perfetto.py
@@ -178,7 +178,8 @@ def build_otel_span_tree(otel_spans):
for span in otel_spans:
parent_span_id = span.get("parentSpanId", "")
- if span.get("attributes", {}).get("module") == "sglang::request":
+ module_name = span.get("attributes", {}).get("module", "")
+ if module_name == "sglang::request" or module_name == "sglang::mooncake":
root_spans.append(span)
elif parent_span_id in span_id_map:
parent_span = span_id_map[parent_span_id]
diff --git a/test/registered/observability/test_tracing_disaggregation.py b/test/registered/observability/test_tracing_disaggregation.py
index f38d513f3..3d8c478e5 100644
--- a/test/registered/observability/test_tracing_disaggregation.py
+++ b/test/registered/observability/test_tracing_disaggregation.py
@@ -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)}",
)
diff --git a/test/registered/unit/observability/test_trace.py b/test/registered/unit/observability/test_trace.py
index 18620e4bf..e26d9bab0 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 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)