feat: add OpenTelemetry tracing to DiffGenerator (#21254)

This commit is contained in:
Jie Hao
2026-04-23 09:25:23 -07:00
committed by GitHub
parent 76e4c5a1f8
commit 86ed0680d7
19 changed files with 978 additions and 259 deletions
@@ -8,6 +8,7 @@ All transfer, compute, and event-loop logic for disaggregated roles
from __future__ import annotations
import contextlib
import dataclasses
import json
import logging
@@ -49,6 +50,8 @@ from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket
from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice
from sglang.srt.observability.trace import TraceReqContext
if TYPE_CHECKING:
from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler
@@ -85,6 +88,12 @@ _EXCLUDE_FIELDS = frozenset(
"step_index",
"prompt_template",
"max_sequence_length",
# trace_ctx holds live OTel SDK objects that aren't JSON-serializable.
# We propagate tracing across the JSON hop via a separate, JSON-safe
# ``_trace_state`` scalar field built from ``TraceReqContext.__getstate__``
# (same W3C carrier SRT relies on for pickle transport) and rebuild it
# on the receiver in ``_build_disagg_req``.
"trace_ctx",
}
)
@@ -235,6 +244,18 @@ def extract_transfer_fields(req) -> tuple[dict, dict]:
_sz,
)
# Propagate OTel trace context over the JSON hop. TraceReqContext.__getstate__
# reduces the live context to a JSON-safe dict (W3C traceparent/tracestate in
# root_span_context). Receiver rebuilds via __setstate__ in _build_disagg_req.
trace_ctx = getattr(req, "trace_ctx", None)
if trace_ctx is not None and getattr(trace_ctx, "tracing_enable", False):
try:
trace_state = trace_ctx.__getstate__()
if trace_state and trace_state.get("tracing_enable"):
scalar_fields["_trace_state"] = trace_state
except Exception as e:
logger.debug("Failed to export trace state: %s", e)
return tensor_fields, scalar_fields
@@ -1235,12 +1256,14 @@ class SchedulerDisaggMixin:
extra_kwargs["mu"] = mu
scheduler_mod.set_timesteps(num_steps, device=device, **extra_kwargs)
self.worker.execute_forward([req], return_req=True)
with self._disagg_trace_dispatch(req):
self.worker.execute_forward([req], return_req=True)
elif self._disagg_role == RoleType.DECODER:
req.save_output = False
req.return_file_paths_only = False
self.worker.execute_forward([req])
with self._disagg_trace_dispatch(req):
self.worker.execute_forward([req])
def _build_disagg_req(self: Scheduler, scalar_fields: dict, tensors: dict) -> Req:
"""Reconstruct a Req from transfer scalar fields and loaded GPU tensors.
@@ -1248,6 +1271,10 @@ class SchedulerDisaggMixin:
Initializes all dataclass field defaults first, then overlays
scalar and tensor fields from the transfer message.
"""
# Pop _trace_state before the generic setattr loop so it doesn't land
# on the Req as a stray attribute.
trace_state = scalar_fields.pop("_trace_state", None)
req = object.__new__(Req)
# Initialize all dataclass fields with their defaults
for f in dataclasses.fields(Req):
@@ -1272,9 +1299,43 @@ class SchedulerDisaggMixin:
gen = torch.Generator(device="cpu")
gen.manual_seed(int(seed))
req.generator = gen
# Rebuild trace_ctx from the propagated __getstate__ dict so this role's
# spans nest under the sender's trace (same mechanism SRT uses via pickle).
if trace_state and trace_state.get("tracing_enable"):
try:
ctx = object.__new__(TraceReqContext)
ctx.__setstate__(trace_state)
req.trace_ctx = ctx
except Exception as e:
logger.debug("Failed to rebuild trace_ctx from _trace_state: %s", e)
req.validate()
return req
@contextlib.contextmanager
def _disagg_trace_dispatch(self: Scheduler, req: Req):
"""Wrap a disagg role's worker.execute_forward in the tracing lifecycle.
Mirrors the monolithic path in ``scheduler._handle_generation``: rebuild
the thread context under the (potentially remote) root_span_context that
was propagated in via ``_trace_state`` / pickle, then emit a
``scheduler_dispatch`` span for this role with ``thread_finish_flag``
so the thread span closes when compute returns. If tracing is disabled
(TraceNullContext), everything is a no-op.
"""
ctx = getattr(req, "trace_ctx", None)
if ctx is None:
yield
return
# Disagg receive (__setstate__) and compute may run on different
# threads (e.g. recv-prefetch vs scheduler main). Align the ctx's pid
# with the current compute thread so __create_thread_context's
# threads_info lookup resolves via the local registration.
if getattr(ctx, "tracing_enable", False):
ctx.pid = threading.get_native_id()
ctx.rebuild_thread_context()
with trace_slice(ctx, DiffStage.SCHEDULER_DISPATCH, thread_finish_flag=True):
yield
def _disagg_denoiser_compute(
self: Scheduler, req: Req, request_id: str, role_name: str
) -> None:
@@ -1285,7 +1346,8 @@ class SchedulerDisaggMixin:
"""
# Run denoising
start_time = time.monotonic()
result = self.worker.execute_forward([req], return_req=True)
with self._disagg_trace_dispatch(req):
result = self.worker.execute_forward([req], return_req=True)
duration_s = time.monotonic() - start_time
if not isinstance(result, Req):
@@ -1375,7 +1437,8 @@ class SchedulerDisaggMixin:
req.return_file_paths_only = False
start_time = time.monotonic()
output_batch = self.worker.execute_forward([req])
with self._disagg_trace_dispatch(req):
output_batch = self.worker.execute_forward([req])
duration_s = time.monotonic() - start_time
# Send result as raw ZMQ frames (no TRANSFER_MAGIC prefix).
@@ -1424,7 +1487,8 @@ class SchedulerDisaggMixin:
self._disagg_metrics.record_request_start(request_id)
# Run encoder stages
req_result = self.worker.execute_forward(reqs, return_req=True)
with self._disagg_trace_dispatch(req):
req_result = self.worker.execute_forward(reqs, return_req=True)
if not isinstance(req_result, Req):
# Error — send error via scalar fields (rank 0 only)
@@ -41,6 +41,8 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
log_batch_completion,
log_generation_timer,
)
from sglang.multimodal_gen.runtime.utils.trace_wrapper import trace_req
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
logger = init_logger(__name__)
@@ -119,6 +121,10 @@ class DiffGenerator:
instance = cls(
server_args=server_args,
)
if server_args.enable_trace:
process_tracing_init(server_args.otlp_traces_endpoint, "sglang-diffusion")
trace_set_thread_info("DiffGenerator")
logger.info(f"Local mode: {local_mode}")
if local_mode:
instance.local_scheduler_process = instance._start_local_server_if_needed()
@@ -176,6 +182,7 @@ class DiffGenerator:
def generate(
self,
sampling_params_kwargs: dict | None = None,
external_trace_header: dict[str, str] | None = None,
) -> GenerationResult | list[GenerationResult] | None:
"""Generate image(s)/video(s) based on the given prompt(s).
@@ -217,6 +224,7 @@ class DiffGenerator:
req = prepare_request(
server_args=self.server_args,
sampling_params=sampling_params,
external_trace_header=external_trace_header,
)
requests.append(req)
@@ -227,7 +235,7 @@ class DiffGenerator:
# TODO: send batch when supported
for request_idx, req in enumerate(requests):
try:
with log_generation_timer(
with trace_req(req.trace_ctx), log_generation_timer(
logger, req.prompt, request_idx + 1, len(requests)
) as timer:
output_batch = self._send_to_scheduler_and_wait_for_response([req])
@@ -6,7 +6,16 @@ import os
import time
from typing import List, Optional
from fastapi import APIRouter, File, Form, HTTPException, Path, Query, UploadFile
from fastapi import (
APIRouter,
File,
Form,
HTTPException,
Path,
Query,
Request,
UploadFile,
)
from fastapi.responses import FileResponse
from sglang.multimodal_gen.configs.sample.sampling_params import generate_request_id
@@ -31,6 +40,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBa
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.observability.trace import extract_trace_headers
router = APIRouter(prefix="/v1/images", tags=["images"])
logger = init_logger(__name__)
@@ -118,6 +128,7 @@ def _build_image_response_kwargs(
@router.post("/generations", response_model=ImageResponse)
async def generations(
request: ImageGenerationsRequest,
raw_request: Request,
):
request_id = generate_request_id()
server_args = get_global_server_args()
@@ -148,9 +159,11 @@ async def generations(
perf_dump_path=request.perf_dump_path,
use_pe=_get_extra_field(request, "use_pe"),
)
trace_headers = extract_trace_headers(raw_request.headers)
batch = prepare_request(
server_args=server_args,
sampling_params=sampling,
external_trace_header=trace_headers,
)
# Add diffusers_kwargs if provided
if request.diffusers_kwargs:
@@ -199,6 +212,7 @@ async def generations(
@router.post("/edits", response_model=ImageResponse)
async def edits(
raw_request: Request,
image: Optional[List[UploadFile]] = File(None),
image_array: Optional[List[UploadFile]] = File(None, alias="image[]"),
url: Optional[List[str]] = Form(None),
@@ -284,9 +298,11 @@ async def edits(
upscaling_model_path=upscaling_model_path,
upscaling_scale=upscaling_scale,
)
trace_headers = extract_trace_headers(raw_request.headers)
batch = prepare_request(
server_args=server_args,
sampling_params=sampling,
external_trace_header=trace_headers,
)
save_file_path_list, result = await process_generation_batch(
async_scheduler_client, batch
@@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
log_batch_completion,
log_generation_timer,
)
from sglang.multimodal_gen.runtime.utils.trace_wrapper import trace_req
# re-export LoRA protocol types for backward compatibility
__all__ = [
@@ -324,7 +325,7 @@ async def process_generation_batch(
batch,
) -> tuple[list[str], OutputBatch]:
total_start_time = time.perf_counter()
with log_generation_timer(logger, batch.prompt):
with trace_req(batch.trace_ctx), log_generation_timer(logger, batch.prompt):
result = await scheduler_client.forward([batch])
if result.output is None and result.output_file_paths is None:
@@ -44,6 +44,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.observability.trace import extract_trace_headers
logger = init_logger(__name__)
router = APIRouter(prefix="/v1/videos", tags=["videos"])
@@ -349,9 +350,11 @@ async def create_video(
await VIDEO_STORE.upsert(request_id, job)
# Build Req for scheduler
trace_headers = extract_trace_headers(request.headers)
batch = prepare_request(
server_args=server_args,
sampling_params=sampling_params,
external_trace_header=trace_headers,
)
# Add diffusers_kwargs if provided
if req.diffusers_kwargs:
@@ -36,6 +36,7 @@ from sglang.multimodal_gen.configs.sample.sampling_params import (
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import CYAN, RESET, init_logger
from sglang.srt.observability.trace import TraceReqContext
logger = init_logger(__name__)
@@ -288,6 +289,7 @@ def _maybe_mux_audio_into_mp4(
def prepare_request(
server_args: ServerArgs,
sampling_params: SamplingParams,
external_trace_header: dict[str, str] | None = None,
) -> Req:
"""
Create a Req object with sampling_params as a parameter.
@@ -310,6 +312,15 @@ def prepare_request(
f"Height and width must be positive, got height={req.height}, width={req.width}"
)
if server_args.enable_trace:
trace_ctx = TraceReqContext(
rid=sampling_params.request_id,
module_name="diffusion",
external_trace_header=external_trace_header,
)
trace_ctx.trace_req_start()
req.trace_ctx = trace_ctx
return req
@@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.server_args import (
)
from sglang.multimodal_gen.runtime.utils.common import is_port_available
from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
def _find_available_port(
@@ -441,6 +442,10 @@ def _run_disagg_role_process(
def launch_http_server_only(server_args):
if server_args.enable_trace:
process_tracing_init(server_args.otlp_traces_endpoint, "sglang-diffusion")
trace_set_thread_info("DiffHTTPServer")
# set for endpoints to access global_server_args
set_global_server_args(server_args)
app = create_app(server_args)
@@ -58,6 +58,8 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import (
PerformanceLogger,
capture_memory_snapshot,
)
from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.utils.network import NetworkAddress
logger = init_logger(__name__)
@@ -234,7 +236,8 @@ class GPUWorker:
req.metrics.record_memory_snapshot("before_forward", baseline_snapshot)
req.log(server_args=self.server_args)
result = self.pipeline.forward(req, self.server_args)
with trace_slice(req.trace_ctx, DiffStage.GPU_FORWARD):
result = self.pipeline.forward(req, self.server_args)
# For disagg roles, return raw Req to let the caller handle
# the role-to-role tensor transfer before OutputBatch conversion.
@@ -530,6 +533,10 @@ def run_scheduler_process(
elif current_platform.is_musa():
set_musa_arch()
if server_args.enable_trace:
process_tracing_init(server_args.otlp_traces_endpoint, "sglang-diffusion")
trace_set_thread_info(f"DiffWorker_rank{rank}")
port_args = PortArgs.from_server_args(server_args)
# start the scheduler event loop
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.server_args import (
from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket
from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj
from sglang.multimodal_gen.runtime.utils.logging_utils import GREEN, RESET, init_logger
from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice
logger = init_logger(__name__)
@@ -205,7 +206,16 @@ class Scheduler(SchedulerDisaggMixin):
else:
logger.info("Processing warmup req...")
return self.worker.execute_forward(reqs)
# Diffusion dispatches one generation request at a time, so reqs[0]
# always carries the trace context for the entire batch.
req = reqs[0]
req.trace_ctx.rebuild_thread_context()
with trace_slice(
req.trace_ctx,
DiffStage.SCHEDULER_DISPATCH,
thread_finish_flag=True,
):
return self.worker.execute_forward(reqs)
def return_result(
self,
@@ -16,7 +16,7 @@ import os
import pprint
from copy import deepcopy
from dataclasses import MISSING, asdict, dataclass, field, fields
from typing import Any, Optional
from typing import Any, Optional, Union
import PIL.Image
import torch
@@ -32,6 +32,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
)
from sglang.multimodal_gen.runtime.utils.perf_logger import RequestMetrics
from sglang.multimodal_gen.utils import align_to
from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
logger = init_logger(__name__)
@@ -161,6 +162,11 @@ class Req:
# stage logging
metrics: Optional["RequestMetrics"] = None
# tracing context (TraceReqContext or TraceNullContext)
trace_ctx: Union[TraceReqContext, TraceNullContext] = field(
default_factory=TraceNullContext
)
# results
output: torch.Tensor | None = None
audio: torch.Tensor | None = None
@@ -282,6 +282,10 @@ class ServerArgs(DisaggArgsMixin):
log_level: str = "info"
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
# Tracing
enable_trace: bool = False
otlp_traces_endpoint: str = "localhost:4317"
# get_role_parallelism, derive_pool_*_endpoint — from DisaggArgsMixin
@property
@@ -1095,6 +1099,20 @@ class ServerArgs(DisaggArgsMixin):
default=ServerArgs.log_level,
help="The logging level of all loggers.",
)
# Tracing
parser.add_argument(
"--enable-trace",
action="store_true",
default=False,
help="Enable OpenTelemetry tracing.",
)
parser.add_argument(
"--otlp-traces-endpoint",
type=str,
default=ServerArgs.otlp_traces_endpoint,
help="OTLP collector endpoint when --enable-trace is set. Format: <host>:<port>",
)
parser.add_argument(
"--uvicorn-access-log-exclude-prefixes",
type=str,
@@ -0,0 +1,57 @@
"""Context-manager wrappers around sglang.srt.observability.trace for diffusion tracing.
All tracing helpers for the multimodal_gen subsystem are consolidated here so
that call sites can use simple ``with`` statements instead of manual
start/end bookkeeping.
"""
from __future__ import annotations
from contextlib import contextmanager
from dataclasses import dataclass
@dataclass(frozen=True)
class DiffStageConfig:
"""A named trace stage with a default nesting level."""
stage_name: str
level: int = 0
class DiffStage:
"""Named trace stages for the diffusion pipeline."""
SCHEDULER_DISPATCH = DiffStageConfig("scheduler_dispatch", level=1)
GPU_FORWARD = DiffStageConfig("gpu_forward", level=2)
@contextmanager
def trace_req(trace_ctx):
"""Ensure ``trace_req_finish()`` is called when a request scope exits.
Usage::
with trace_req(batch.trace_ctx):
...
"""
try:
yield trace_ctx
finally:
trace_ctx.trace_req_finish()
@contextmanager
def trace_slice(trace_ctx, stage: DiffStageConfig, **kwargs):
"""Context manager for a single trace slice (span).
Usage::
with trace_slice(req.trace_ctx, DiffStage.GPU_FORWARD):
result = pipeline.forward(req, server_args)
"""
trace_ctx.trace_slice_start(stage.stage_name, level=stage.level)
try:
yield trace_ctx
finally:
trace_ctx.trace_slice_end(stage.stage_name, level=stage.level, **kwargs)
@@ -82,6 +82,7 @@ STANDALONE_FILES = {
"1-gpu": [
"../cli/test_generate_t2i_perf.py",
"test_update_weights_from_disk.py",
"test_tracing.py",
],
"2-gpu": [
"test_disagg_server.py",
@@ -95,6 +96,7 @@ STANDALONE_FILE_EST_TIMES = {
"1-gpu": {
"../cli/test_generate_t2i_perf.py": 240.0,
"test_update_weights_from_disk.py": 480.0,
"test_tracing.py": 120.0,
},
"2-gpu": {
# Two disagg clusters × (~3 min startup + ~1 min generate) ≈ 8 min.
@@ -202,6 +202,10 @@ class DisaggCluster:
def _launch_server_head(self) -> None:
log = _LOG_DIR / f"disagg_{self.name}_server.log"
self._logs["server"] = log
# Role processes register their transfer work_endpoint with the
# derived value ``tcp://0.0.0.0:<port>`` (see disagg_args.py). The
# server head must advertise the same literal so ``_handle_register``'s
# endpoint_to_idx exact-string match succeeds.
cmd = [
"sglang",
"serve",
@@ -210,11 +214,11 @@ class DisaggCluster:
"--disagg-role",
"server",
"--encoder-urls",
f"tcp://{HOST}:{self._role_ports['encoder']}",
f"tcp://0.0.0.0:{self._role_ports['encoder']}",
"--denoiser-urls",
f"tcp://{HOST}:{self._role_ports['denoiser']}",
f"tcp://0.0.0.0:{self._role_ports['denoiser']}",
"--decoder-urls",
f"tcp://{HOST}:{self._role_ports['decoder']}",
f"tcp://0.0.0.0:{self._role_ports['decoder']}",
"--scheduler-port",
str(self.base_port),
"--port",
@@ -225,6 +229,7 @@ class DisaggCluster:
"120",
"--log-level",
"info",
*self.extra_role_args.get("server", []),
]
self._start_proc(cmd, log)
try:
@@ -384,5 +389,171 @@ class TestDisaggZImage2RankDenoiser(_DisaggTestBase):
self.assertGreater(len(img), 1_000, f"image too small: {len(img)} bytes")
# ---------------------------------------------------------------------------
# Disagg + OTel tracing
# ---------------------------------------------------------------------------
def _generate_image_with_traceparent(
api_port: int, model: str, trace_id_hex: str, span_id_hex: str
) -> tuple[int, bytes]:
"""Same as :func:`_generate_image` but seeds a known W3C traceparent.
Returns ``(status_code, image_bytes)``. Kept separate so the tracing test
can tolerate non-200 responses while still reporting useful diagnostics.
"""
traceparent = f"00-{trace_id_hex}-{span_id_hex}-01"
resp = requests.post(
f"http://{HOST}:{api_port}/v1/images/generations",
headers={"traceparent": traceparent},
json={
"model": model,
"prompt": "A sunset over mountains",
"n": 1,
"size": "1024x1024",
"response_format": "b64_json",
},
timeout=600,
)
if resp.status_code != 200:
return resp.status_code, b""
return resp.status_code, base64.b64decode(resp.json()["data"][0]["b64_json"])
def _as_hex(v) -> str:
"""OTLP span trace_id/span_id/parent_span_id come back as raw bytes over
gRPC and as hex strings over HTTP; normalize to lowercase hex."""
if isinstance(v, (bytes, bytearray)):
return v.hex()
if isinstance(v, str):
return v.lower()
return ""
class TestDisaggZImageTracing(_DisaggTestBase):
"""End-to-end verification of OTel trace propagation across disagg roles.
Spins up the same 1-rank cluster as :class:`TestDisaggZImage1Rank` with
``--enable-trace`` wired to an in-process OTLP collector on every role and
the server head, sends one image-generation request with a controlled
``traceparent``, and asserts the server head plus all three role worker
processes emit per-role ``scheduler_dispatch``/``gpu_forward`` spans under
the same trace_id. This is the regression guard for trace-context
propagation over the encoder→denoiser→decoder JSON hops.
"""
cluster_name = "zimage_trace"
required_gpus = 2
gpu_layout = {
"encoder": [0],
"denoiser": [1],
"decoder": [0],
}
# Populated in setUpClass so the collector port is known before
# DisaggCluster launches.
collector = None
collector_port: int = 0
@classmethod
def setUpClass(cls) -> None:
# Fast batch-span-processor flush so the test doesn't wait for the
# default 5s schedule. Must be set before sglang imports OTel.
os.environ.setdefault("SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS", "50")
os.environ.setdefault("SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE", "4")
from sglang.test.otel_collector import LightweightOtlpCollector
cls.collector_port = find_free_port(HOST)
cls.collector = LightweightOtlpCollector(port=cls.collector_port)
cls.collector.start()
trace_args = [
"--enable-trace",
"--otlp-traces-endpoint",
f"127.0.0.1:{cls.collector_port}",
]
cls.extra_role_args = {
"encoder": list(trace_args),
"denoiser": list(trace_args),
"decoder": list(trace_args),
"server": list(trace_args),
}
# If super().setUpClass() raises, CustomTestCase's safe-setUpClass
# wrapper will invoke tearDownClass, which stops the collector.
super().setUpClass()
@classmethod
def tearDownClass(cls) -> None:
try:
super().tearDownClass()
finally:
if cls.collector is not None:
cls.collector.stop()
cls.collector = None
def test_disagg_spans_share_trace_id(self) -> None:
assert self.cluster is not None
assert self.collector is not None
trace_id = os.urandom(16).hex()
span_id = os.urandom(8).hex()
# Warmup was sent (without traceparent) by DisaggCluster.__enter__;
# clear those spans so the assertions only consider this request.
self.collector.clear()
status, img = _generate_image_with_traceparent(
self.cluster.api_port, self.model, trace_id, span_id
)
self.assertEqual(status, 200, "request did not complete cleanly")
self.assertGreater(len(img), 1_000, f"image too small: {len(img)} bytes")
# Spans flush asynchronously from each role's BatchSpanProcessor. Poll
# briefly until we see the expected shape.
deadline = time.time() + 30
spans = []
while time.time() < deadline:
spans = [
s for s in self.collector.get_spans() if _as_hex(s.trace_id) == trace_id
]
# Expect: root Req span + >=3 scheduler_dispatch + >=3 gpu_forward
n_dispatch = sum(1 for s in spans if s.name == "scheduler_dispatch")
n_forward = sum(1 for s in spans if s.name == "gpu_forward")
if n_dispatch >= 3 and n_forward >= 3:
break
time.sleep(1)
names = [s.name for s in spans]
self.assertGreaterEqual(
sum(1 for n in names if n == "scheduler_dispatch"),
3,
f"expected >=3 scheduler_dispatch spans (one per disagg role), "
f"got names={names!r}",
)
self.assertGreaterEqual(
sum(1 for n in names if n == "gpu_forward"),
3,
f"expected >=3 gpu_forward spans (one per disagg role), "
f"got names={names!r}",
)
# All spans we saw must share the propagated trace_id. This is the
# actual regression guard for this PR: it proves the W3C carrier
# survives encoder→denoiser→decoder JSON hops (via ``_trace_state``).
# The HTTP-level carrier extraction (root Req parented under the
# client's span_id) is already covered by ``test_tracing.py`` in
# monolithic mode and asserting it here is flaky — the server head's
# BatchSpanProcessor may not flush the Req span before role spans
# reach the collector, since the role spans close first.
trace_ids = {_as_hex(s.trace_id) for s in spans}
self.assertEqual(
trace_ids,
{trace_id},
f"spans split across multiple traces: {trace_ids}",
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,155 @@
"""Integration test for OpenTelemetry tracing in the diffusion pipeline.
Spins up a lightweight in-process OTLP collector, launches a diffusion server
with ``--enable-trace``, sends an image-generation request with a
``traceparent`` header, and asserts that the expected spans
(``scheduler_dispatch``, ``gpu_forward``) are exported.
"""
import os
# Configure OTLP exporter for faster test execution.
# Must be set before importing any sglang trace module.
os.environ.setdefault("SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS", "50")
os.environ.setdefault("SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE", "4")
import logging
import time
import pytest
import requests
from sglang.multimodal_gen.test.server.test_server_utils import ServerManager
from sglang.multimodal_gen.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST
from sglang.test.otel_collector import LightweightOtlpCollector
logger = logging.getLogger(__name__)
# Expected diffusion trace span names (from DiffStage in trace_wrapper.py)
EXPECTED_DIFF_SPANS = ["scheduler_dispatch", "gpu_forward"]
COLLECTOR_PORT = 4317
SERVER_PORT = 39812
@pytest.fixture(scope="module")
def tracing_env():
"""Start the OTLP collector and diffusion server once for all tests."""
collector = LightweightOtlpCollector(port=COLLECTOR_PORT)
collector.start()
time.sleep(0.3)
mgr = ServerManager(
model=DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
port=SERVER_PORT,
extra_args=f"--enable-trace --otlp-traces-endpoint 127.0.0.1:{COLLECTOR_PORT}",
)
ctx = mgr.start()
# Clear any warmup spans
time.sleep(2)
collector.clear()
yield collector, ctx
ctx.cleanup()
collector.stop()
def _generate_image(headers=None):
"""Send a single image-generation request."""
resp = requests.post(
f"http://127.0.0.1:{SERVER_PORT}/v1/images/generations",
json={
"model": DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
"prompt": "A white cat",
"size": "256x256",
"n": 1,
},
headers=headers or {},
timeout=300,
)
assert resp.status_code == 200, f"Generation failed: {resp.text}"
return resp
def _wait_for_spans(collector, required_names=None, min_count=1, timeout=30):
"""Wait until collector has the required span names (or at least ``min_count`` spans)."""
deadline = time.time() + timeout
while time.time() < deadline:
if required_names:
if all(collector.has_span(n) for n in required_names):
return
elif collector.count_spans() >= min_count:
return
time.sleep(0.5)
def test_spans_exported(tracing_env):
"""After a generation request the expected diffusion spans appear."""
collector, _ = tracing_env
collector.clear()
# W3C Trace Context traceparent header
_generate_image(
headers={
"traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"
}
)
_wait_for_spans(collector, required_names=EXPECTED_DIFF_SPANS)
span_names = collector.get_span_names()
for expected in EXPECTED_DIFF_SPANS:
assert (
expected in span_names
), f"Missing span '{expected}'. Collected: {sorted(span_names)}"
def test_spans_without_traceparent(tracing_env):
"""Requests without a traceparent header still produce spans as a new
root trace (not linked to any upstream)."""
collector, _ = tracing_env
collector.clear()
_generate_image()
_wait_for_spans(collector, required_names=EXPECTED_DIFF_SPANS)
span_names = collector.get_span_names()
for expected in EXPECTED_DIFF_SPANS:
assert (
expected in span_names
), f"Missing span '{expected}'. Collected: {sorted(span_names)}"
def test_batch_requests(tracing_env):
"""Multiple requests each produce their own set of spans."""
collector, _ = tracing_env
collector.clear()
batch_size = 3
for i in range(batch_size):
# Each request gets a unique trace-id
trace_id = f"0af7651916cd43dd8448eb211c8031{i:02x}"
_generate_image(headers={"traceparent": f"00-{trace_id}-b7ad6b7169203331-01"})
# Wait until all scheduler_dispatch spans have arrived (they come from a
# separate process so may lag behind gpu_forward).
deadline = time.time() + 60
while time.time() < deadline:
if all(
len(collector.get_spans_by_name(n)) >= batch_size
for n in EXPECTED_DIFF_SPANS
):
break
time.sleep(0.5)
for span_name in EXPECTED_DIFF_SPANS:
matching = collector.get_spans_by_name(span_name)
assert len(matching) >= batch_size, (
f"Expected at least {batch_size} '{span_name}' spans, "
f"got {len(matching)}"
)
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])
@@ -0,0 +1,166 @@
"""Unit tests for OTel trace-context propagation across the diffusion disagg
JSON hop (encoder -> denoiser, denoiser -> decoder).
These exercise the serialization contract only (no GPUs, no server, no OTLP
collector required):
- ``extract_transfer_fields`` emits a JSON-safe ``_trace_state`` (W3C carrier)
when tracing is enabled, and omits it when tracing is disabled. It never
serializes the live ``TraceReqContext`` object itself.
- The ``_trace_state`` payload round-trips through ``codec.pack_tensors``
(the same ``json.dumps`` path the RDMA metadata frame uses).
- ``TraceReqContext.__setstate__`` reconstructs a live, ``is_copy=True``
context whose ``root_span_context`` is an OTel ``Context`` object.
- ``_build_disagg_req`` pops ``_trace_state`` and installs a rebuilt
``TraceReqContext`` on the receiver-side Req.
"""
from __future__ import annotations
import json
import unittest
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
SchedulerDisaggMixin,
extract_transfer_fields,
)
from sglang.multimodal_gen.runtime.disaggregation.transport.codec import pack_tensors
from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.srt.observability import trace as srt_trace
from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
try:
from opentelemetry import propagate as otel_propagate
from opentelemetry import trace as otel_trace
from opentelemetry.sdk.trace import TracerProvider
_OTEL_AVAILABLE = True
except ImportError:
_OTEL_AVAILABLE = False
_OTEL_BOOTSTRAPPED = False
def _enable_minimal_otel() -> None:
"""Bootstrap just enough OTel state for TraceReqContext to produce real
spans. Idempotent — the TracerProvider can only be set once per process."""
global _OTEL_BOOTSTRAPPED
if not _OTEL_BOOTSTRAPPED:
otel_trace.set_tracer_provider(TracerProvider())
_OTEL_BOOTSTRAPPED = True
srt_trace.opentelemetry_initialized = True
srt_trace.tracer = otel_trace.get_tracer("test-diffusion-disagg")
srt_trace.trace_set_thread_info("TestThread")
def _traceparent_from(ctx) -> str | None:
"""Re-inject a W3C carrier from an OTel Context and return the traceparent.
Used to assert that a carrier round-trip preserves trace_id/span_id, which
is the actual correctness property (OTel's Context is a dict subclass, so
``isinstance(ctx, dict)`` isn't useful).
"""
carrier: dict = {}
otel_propagate.inject(carrier, ctx)
return carrier.get("traceparent")
def _roundtrip_scalar_fields(scalar_fields: dict) -> dict:
"""Run the actual RDMA metadata codec path: pack -> json bytes -> decode."""
metadata_bytes, _ = pack_tensors({}, scalar_fields)
decoded = json.loads(metadata_bytes.decode("utf-8"))
return decoded["scalar_fields"]
class TestDisaggTracePropagation(unittest.TestCase):
def test_tracing_disabled_omits_trace_state(self):
"""With a default TraceNullContext Req, no _trace_state is emitted and
the JSON codec does not encounter any live OTel objects."""
req = Req(request_id="test-off", prompt="x")
self.assertIsInstance(req.trace_ctx, TraceNullContext)
_, scalar_fields = extract_transfer_fields(req)
self.assertNotIn("_trace_state", scalar_fields)
# trace_ctx must never ride the JSON scalar path.
self.assertNotIn("trace_ctx", scalar_fields)
# json.dumps must succeed (this is the path that pre-fix crashed).
_roundtrip_scalar_fields(scalar_fields)
@unittest.skipUnless(_OTEL_AVAILABLE, "opentelemetry SDK not installed")
def test_tracing_enabled_state_roundtrip(self):
"""Sender emits a W3C-carrier _trace_state, it round-trips through
json encode/decode, and __setstate__ reconstructs a live is_copy=True
TraceReqContext with an OTel Context (not the raw dict)."""
_enable_minimal_otel()
ctx = TraceReqContext(rid="test-on", role="server", module_name="request")
ctx.trace_req_start()
self.assertTrue(ctx.tracing_enable)
self.assertFalse(ctx.is_copy)
req = Req(request_id="test-on", prompt="x")
req.trace_ctx = ctx
_, scalar_fields = extract_transfer_fields(req)
self.assertNotIn("trace_ctx", scalar_fields)
self.assertIn("_trace_state", scalar_fields)
state = scalar_fields["_trace_state"]
self.assertTrue(state.get("tracing_enable"))
# W3C carrier must be present so downstream roles can nest spans.
self.assertIn("traceparent", state.get("root_span_context", {}))
decoded = _roundtrip_scalar_fields(scalar_fields)
self.assertEqual(decoded["_trace_state"], state)
rebuilt = object.__new__(TraceReqContext)
rebuilt.__setstate__(decoded["_trace_state"])
self.assertTrue(rebuilt.tracing_enable)
self.assertTrue(rebuilt.is_copy)
# The sender's traceparent must survive into the rebuilt Context so
# downstream role spans nest under the original trace_id.
self.assertEqual(
_traceparent_from(rebuilt.root_span_context),
state["root_span_context"]["traceparent"],
)
@unittest.skipUnless(_OTEL_AVAILABLE, "opentelemetry SDK not installed")
def test_build_disagg_req_installs_rebuilt_ctx(self):
"""_build_disagg_req pops _trace_state from scalar_fields and installs
a live TraceReqContext on the rebuilt Req; the key does not leak onto
the Req as a stray attribute."""
_enable_minimal_otel()
ctx = TraceReqContext(rid="test-brq", role="server", module_name="request")
ctx.trace_req_start()
req = Req(request_id="test-brq", prompt="x")
req.trace_ctx = ctx
_, scalar_fields = extract_transfer_fields(req)
self.assertIn("_trace_state", scalar_fields)
# _build_disagg_req is an instance method but its body does not touch
# ``self``; call via __func__ to avoid needing a real Scheduler.
rebuilt = SchedulerDisaggMixin._build_disagg_req(None, dict(scalar_fields), {})
self.assertIsInstance(rebuilt.trace_ctx, TraceReqContext)
self.assertTrue(rebuilt.trace_ctx.tracing_enable)
self.assertTrue(rebuilt.trace_ctx.is_copy)
self.assertFalse(hasattr(rebuilt, "_trace_state"))
@unittest.skipUnless(_OTEL_AVAILABLE, "opentelemetry SDK not installed")
def test_build_disagg_req_falls_back_when_tracing_off(self):
"""If the sender's context is a TraceNullContext, the receiver's Req
keeps its default TraceNullContext (no _trace_state to apply)."""
req = Req(request_id="test-brq-off", prompt="x")
self.assertIsInstance(req.trace_ctx, TraceNullContext)
_, scalar_fields = extract_transfer_fields(req)
self.assertNotIn("_trace_state", scalar_fields)
rebuilt = SchedulerDisaggMixin._build_disagg_req(None, dict(scalar_fields), {})
self.assertIsInstance(rebuilt.trace_ctx, TraceNullContext)
if __name__ == "__main__":
unittest.main()
+259
View File
@@ -0,0 +1,259 @@
"""Lightweight in-process OTLP collector for tracing tests.
Provides a minimal OTLP collector that receives traces via gRPC (with HTTP
fallback) and stores them in memory for test assertions, eliminating the need
for Docker-based opentelemetry-collector and file I/O.
Usage::
collector = LightweightOtlpCollector(port=4317)
collector.start()
# ... run code that emits traces ...
assert collector.has_span("my_span")
collector.stop()
"""
import json
import logging
import threading
from concurrent import futures
from dataclasses import dataclass, field
from typing import Any, Dict, List, Set
logger = logging.getLogger(__name__)
@dataclass
class Span:
"""Represents a single span extracted from OTLP trace data."""
name: str
trace_id: str = ""
span_id: str = ""
parent_span_id: str = ""
start_time_ns: int = 0
end_time_ns: int = 0
attributes: Dict[str, Any] = field(default_factory=dict)
events: List[Dict[str, Any]] = field(default_factory=list)
class LightweightOtlpCollector:
"""A minimal OTLP collector that stores traces in memory for test assertions.
This replaces the Docker-based opentelemetry-collector for testing purposes.
It listens on a gRPC port for OTLP trace data and stores spans in memory,
allowing tests to verify specific spans based on trace level.
"""
def __init__(self, port: int = 4317):
self.port = port
self._server = None
self._thread = None
self._running = False
self._lock = threading.Lock()
# In-memory storage for collected spans
self._spans: List[Span] = []
self._raw_traces: List[Dict[str, Any]] = []
def _try_grpc_server(self):
"""Try to start gRPC server with full OTLP protocol."""
try:
from grpc import server as grpc_server
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import (
ExportTraceServiceResponse,
)
from opentelemetry.proto.collector.trace.v1.trace_service_pb2_grpc import (
TraceServiceServicer,
add_TraceServiceServicer_to_server,
)
class TraceServicer(TraceServiceServicer):
def __init__(self, collector):
self.collector = collector
def Export(self, request, context):
self.collector._handle_trace_request(request)
return ExportTraceServiceResponse()
self._server = grpc_server(futures.ThreadPoolExecutor(max_workers=4))
add_TraceServiceServicer_to_server(TraceServicer(self), self._server)
self._server.add_insecure_port(f"127.0.0.1:{self.port}")
return True
except ImportError:
logger.warning("Full gRPC OTLP not available, using HTTP fallback")
return False
def _handle_trace_request(self, request):
"""Handle incoming trace request and extract spans to memory."""
with self._lock:
try:
trace_data = self._protobuf_to_dict(request)
self._raw_traces.append(trace_data)
# Extract spans from the trace data
self._extract_spans(trace_data)
except Exception as e:
logger.error(f"Failed to process trace: {e}")
def _protobuf_to_dict(self, proto_obj) -> Dict[str, Any]:
"""Convert protobuf message to dict."""
result = {}
for field, value in proto_obj.ListFields():
if field.message_type:
type_name = type(value).__name__
if "Repeated" in type_name:
result[field.name] = [self._protobuf_to_dict(v) for v in value]
else:
result[field.name] = self._protobuf_to_dict(value)
else:
result[field.name] = value
return result
def _extract_spans(self, trace_data: Dict[str, Any]):
"""Extract Span objects from OTLP trace data structure."""
resource_spans = trace_data.get("resource_spans", [])
for rs in resource_spans:
scope_spans = rs.get("scope_spans", [])
for ss in scope_spans:
spans = ss.get("spans", [])
for span_data in spans:
span = Span(
name=span_data.get("name", ""),
trace_id=span_data.get("trace_id", ""),
span_id=span_data.get("span_id", ""),
parent_span_id=span_data.get("parent_span_id", ""),
start_time_ns=span_data.get("start_time_unix_nano", 0),
end_time_ns=span_data.get("end_time_unix_nano", 0),
attributes=span_data.get("attributes", {}),
events=span_data.get("events", []),
)
self._spans.append(span)
def _http_server_loop(self):
"""Fallback HTTP server for OTLP HTTP protocol."""
from http.server import BaseHTTPRequestHandler, HTTPServer
class OTLPHandler(BaseHTTPRequestHandler):
def __init__(self, request, client_address, server):
self.collector = server.collector
super().__init__(request, client_address, server)
def do_POST(self):
if self.path in ["/v1/traces", "/v1/traces/"]:
content_length = int(self.headers.get("Content-Length", 0))
body = self.rfile.read(content_length)
try:
data = json.loads(body)
with self.collector._lock:
self.collector._raw_traces.append(data)
self.collector._extract_spans_http(data)
self.send_response(200)
self.end_headers()
except Exception as e:
logger.error(f"HTTP trace handling error: {e}")
self.send_response(500)
self.end_headers()
else:
self.send_response(404)
self.end_headers()
def log_message(self, format, *args):
pass # Suppress HTTP server logs
class CollectorHTTPServer(HTTPServer):
def __init__(self, server_address, collector):
self.collector = collector
super().__init__(
server_address,
lambda r, a, s: OTLPHandler(r, a, s),
)
server = CollectorHTTPServer(("127.0.0.1", 4318), self)
server.timeout = 0.5
while self._running:
server.handle_request()
def _extract_spans_http(self, data: Dict[str, Any]):
"""Extract Span objects from OTLP HTTP JSON format."""
resource_spans = data.get("resourceSpans", [])
for rs in resource_spans:
scope_spans = rs.get("scopeSpans", [])
for ss in scope_spans:
spans = ss.get("spans", [])
for span_data in spans:
span = Span(
name=span_data.get("name", ""),
trace_id=span_data.get("traceId", ""),
span_id=span_data.get("spanId", ""),
parent_span_id=span_data.get("parentSpanId", ""),
start_time_ns=span_data.get("startTimeUnixNano", 0),
end_time_ns=span_data.get("endTimeUnixNano", 0),
attributes=span_data.get("attributes", {}),
events=span_data.get("events", []),
)
self._spans.append(span)
def start(self):
"""Start the collector server."""
self._running = True
self._spans.clear()
self._raw_traces.clear()
if self._try_grpc_server():
self._server.start()
logger.info(f"OTLP gRPC collector started on port {self.port}")
else:
# Fallback to HTTP server in a thread
self._thread = threading.Thread(target=self._http_server_loop, daemon=True)
self._thread.start()
logger.info("OTLP HTTP collector started on port 4318")
def stop(self):
"""Stop the collector server."""
self._running = False
if self._server:
self._server.stop(1)
self._server = None
logger.info("OTLP collector stopped")
# ========================================================================
# Public API for test assertions
# ========================================================================
def get_spans(self) -> List[Span]:
"""Get all collected spans."""
with self._lock:
return list(self._spans)
def get_span_names(self) -> Set[str]:
"""Get all unique span names."""
with self._lock:
return {s.name for s in self._spans}
def has_span(self, name: str) -> bool:
"""Check if a span with the given name exists."""
return name in self.get_span_names()
def has_any_span(self, names: List[str]) -> bool:
"""Check if any of the given span names exist."""
span_names = self.get_span_names()
return any(name in span_names for name in names)
def has_all_spans(self, names: List[str]) -> bool:
"""Check if all of the given span names exist."""
span_names = self.get_span_names()
return all(name in span_names for name in names)
def get_spans_by_name(self, name: str) -> List[Span]:
"""Get all spans with the given name."""
with self._lock:
return [s for s in self._spans if s.name == name]
def count_spans(self) -> int:
"""Get total count of collected spans."""
with self._lock:
return len(self._spans)
def clear(self):
"""Clear all collected spans."""
with self._lock:
self._spans.clear()
self._raw_traces.clear()
+4 -242
View File
@@ -12,15 +12,12 @@ import os
os.environ.setdefault("SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS", "50")
os.environ.setdefault("SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE", "4")
import json
import logging
import multiprocessing as mp
import threading
import time
import unittest
from concurrent import futures
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Set, Union
from dataclasses import dataclass
from typing import List, Optional, Union
import requests
import zmq
@@ -53,245 +50,10 @@ register_cuda_ci(est_time=104, suite="stage-b-test-1-gpu-small")
# ============================================================================
# Lightweight OTLP Collector (replaces Docker-based otel-collector)
# Lightweight OTLP Collector (shared across tracing tests)
# ============================================================================
@dataclass
class Span:
"""Represents a single span extracted from OTLP trace data."""
name: str
trace_id: str = ""
span_id: str = ""
parent_span_id: str = ""
start_time_ns: int = 0
end_time_ns: int = 0
attributes: Dict[str, Any] = field(default_factory=dict)
events: List[Dict[str, Any]] = field(default_factory=list)
class LightweightOtlpCollector:
"""A minimal OTLP collector that stores traces in memory for test assertions.
This replaces the Docker-based opentelemetry-collector for testing purposes.
It listens on a gRPC port for OTLP trace data and stores spans in memory,
allowing tests to verify specific spans based on trace level.
"""
def __init__(self, port: int = 4317):
self.port = port
self._server = None
self._thread = None
self._running = False
self._lock = threading.Lock()
# In-memory storage for collected spans
self._spans: List[Span] = []
self._raw_traces: List[Dict[str, Any]] = []
def _try_grpc_server(self):
"""Try to start gRPC server with full OTLP protocol."""
try:
from grpc import server as grpc_server
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import (
ExportTraceServiceResponse,
)
from opentelemetry.proto.collector.trace.v1.trace_service_pb2_grpc import (
TraceServiceServicer,
add_TraceServiceServicer_to_server,
)
class TraceServicer(TraceServiceServicer):
def __init__(self, collector):
self.collector = collector
def Export(self, request, context):
self.collector._handle_trace_request(request)
return ExportTraceServiceResponse()
self._server = grpc_server(futures.ThreadPoolExecutor(max_workers=4))
add_TraceServiceServicer_to_server(TraceServicer(self), self._server)
self._server.add_insecure_port(f"127.0.0.1:{self.port}")
return True
except ImportError:
logger.warning("Full gRPC OTLP not available, using HTTP fallback")
return False
def _handle_trace_request(self, request):
"""Handle incoming trace request and extract spans to memory."""
with self._lock:
try:
trace_data = self._protobuf_to_dict(request)
self._raw_traces.append(trace_data)
# Extract spans from the trace data
self._extract_spans(trace_data)
except Exception as e:
logger.error(f"Failed to process trace: {e}")
def _protobuf_to_dict(self, proto_obj) -> Dict[str, Any]:
"""Convert protobuf message to dict."""
result = {}
for field, value in proto_obj.ListFields():
if field.message_type:
type_name = type(value).__name__
if "Repeated" in type_name:
result[field.name] = [self._protobuf_to_dict(v) for v in value]
else:
result[field.name] = self._protobuf_to_dict(value)
else:
result[field.name] = value
return result
def _extract_spans(self, trace_data: Dict[str, Any]):
"""Extract Span objects from OTLP trace data structure."""
resource_spans = trace_data.get("resource_spans", [])
for rs in resource_spans:
scope_spans = rs.get("scope_spans", [])
for ss in scope_spans:
spans = ss.get("spans", [])
for span_data in spans:
span = Span(
name=span_data.get("name", ""),
trace_id=span_data.get("trace_id", ""),
span_id=span_data.get("span_id", ""),
parent_span_id=span_data.get("parent_span_id", ""),
start_time_ns=span_data.get("start_time_unix_nano", 0),
end_time_ns=span_data.get("end_time_unix_nano", 0),
attributes=span_data.get("attributes", {}),
events=span_data.get("events", []),
)
self._spans.append(span)
def _http_server_loop(self):
"""Fallback HTTP server for OTLP HTTP protocol."""
from http.server import BaseHTTPRequestHandler, HTTPServer
class OTLPHandler(BaseHTTPRequestHandler):
def __init__(self, request, client_address, server):
self.collector = server.collector
super().__init__(request, client_address, server)
def do_POST(self):
if self.path in ["/v1/traces", "/v1/traces/"]:
content_length = int(self.headers.get("Content-Length", 0))
body = self.rfile.read(content_length)
try:
data = json.loads(body)
with self.collector._lock:
self.collector._raw_traces.append(data)
self.collector._extract_spans_http(data)
self.send_response(200)
self.end_headers()
except Exception as e:
logger.error(f"HTTP trace handling error: {e}")
self.send_response(500)
self.end_headers()
else:
self.send_response(404)
self.end_headers()
def log_message(self, format, *args):
pass # Suppress HTTP server logs
class CollectorHTTPServer(HTTPServer):
def __init__(self, server_address, collector):
self.collector = collector
super().__init__(
server_address,
lambda r, a, s: OTLPHandler(r, a, s),
)
server = CollectorHTTPServer(("127.0.0.1", 4318), self)
server.timeout = 0.5
while self._running:
server.handle_request()
def _extract_spans_http(self, data: Dict[str, Any]):
"""Extract Span objects from OTLP HTTP JSON format."""
resource_spans = data.get("resourceSpans", [])
for rs in resource_spans:
scope_spans = rs.get("scopeSpans", [])
for ss in scope_spans:
spans = ss.get("spans", [])
for span_data in spans:
span = Span(
name=span_data.get("name", ""),
trace_id=span_data.get("traceId", ""),
span_id=span_data.get("spanId", ""),
parent_span_id=span_data.get("parentSpanId", ""),
start_time_ns=span_data.get("startTimeUnixNano", 0),
end_time_ns=span_data.get("endTimeUnixNano", 0),
attributes=span_data.get("attributes", {}),
events=span_data.get("events", []),
)
self._spans.append(span)
def start(self):
"""Start the collector server."""
self._running = True
self._spans.clear()
self._raw_traces.clear()
if self._try_grpc_server():
self._server.start()
logger.info(f"OTLP gRPC collector started on port {self.port}")
else:
# Fallback to HTTP server in a thread
self._thread = threading.Thread(target=self._http_server_loop, daemon=True)
self._thread.start()
logger.info("OTLP HTTP collector started on port 4318")
def stop(self):
"""Stop the collector server."""
self._running = False
if self._server:
self._server.stop(1)
self._server = None
logger.info("OTLP collector stopped")
# ========================================================================
# Public API for test assertions
# ========================================================================
def get_spans(self) -> List[Span]:
"""Get all collected spans."""
with self._lock:
return list(self._spans)
def get_span_names(self) -> Set[str]:
"""Get all unique span names."""
with self._lock:
return {s.name for s in self._spans}
def has_span(self, name: str) -> bool:
"""Check if a span with the given name exists."""
return name in self.get_span_names()
def has_any_span(self, names: List[str]) -> bool:
"""Check if any of the given span names exist."""
span_names = self.get_span_names()
return any(name in span_names for name in names)
def has_all_spans(self, names: List[str]) -> bool:
"""Check if all of the given span names exist."""
span_names = self.get_span_names()
return all(name in span_names for name in names)
def get_spans_by_name(self, name: str) -> List[Span]:
"""Get all spans with the given name."""
with self._lock:
return [s for s in self._spans if s.name == name]
def count_spans(self) -> int:
"""Get total count of collected spans."""
with self._lock:
return len(self._spans)
def clear(self):
"""Clear all collected spans."""
with self._lock:
self._spans.clear()
self._raw_traces.clear()
from sglang.test.otel_collector import LightweightOtlpCollector, Span # noqa: F401
# ============================================================================
# Test Helper Functions
@@ -15,12 +15,10 @@ from urllib.parse import urlparse
import requests
# Import the lightweight collector from the main tracing test module
from test_tracing import LightweightOtlpCollector
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
from sglang.test.otel_collector import LightweightOtlpCollector
from sglang.test.server_fixtures.disaggregation_fixture import get_rdma_devices_args
from sglang.test.test_utils import (
DEFAULT_MODEL_NAME_FOR_TEST,