diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 670d0dedb..d9e650619 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 7e9a60c10..2af07ec83 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -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]) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index 8e6697157..43cb01d75 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index d8a48c7d9..34a29a9ef 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index 9798eff1b..7592d1791 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index a8b210bbd..2b693f38b 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index 3a1bafa69..002710af6 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index ad206f92e..fdb8f102b 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 129d99dfc..48210705d 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index 674cc1934..7f5b332e9 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 65db9dc2f..07f7d3af0 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -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: :", + ) parser.add_argument( "--uvicorn-access-log-exclude-prefixes", type=str, diff --git a/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py b/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py new file mode 100644 index 000000000..e8a0ecb3c --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/trace_wrapper.py @@ -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) diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index c3ce5e0f8..975a027ea 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -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. diff --git a/python/sglang/multimodal_gen/test/server/test_disagg_server.py b/python/sglang/multimodal_gen/test/server/test_disagg_server.py index e696c8f59..a8233a6fd 100755 --- a/python/sglang/multimodal_gen/test/server/test_disagg_server.py +++ b/python/sglang/multimodal_gen/test/server/test_disagg_server.py @@ -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:`` (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() diff --git a/python/sglang/multimodal_gen/test/server/test_tracing.py b/python/sglang/multimodal_gen/test/server/test_tracing.py new file mode 100644 index 000000000..5da984e3b --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/test_tracing.py @@ -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"]) diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py b/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py new file mode 100644 index 000000000..c8d3861ab --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py @@ -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() diff --git a/python/sglang/test/otel_collector.py b/python/sglang/test/otel_collector.py new file mode 100644 index 000000000..ca9358907 --- /dev/null +++ b/python/sglang/test/otel_collector.py @@ -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() diff --git a/test/registered/observability/test_tracing.py b/test/registered/observability/test_tracing.py index c79b325db..f77dc2ea2 100644 --- a/test/registered/observability/test_tracing.py +++ b/test/registered/observability/test_tracing.py @@ -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 diff --git a/test/registered/observability/test_tracing_disaggregation.py b/test/registered/observability/test_tracing_disaggregation.py index f06d2718d..91d1e11ed 100644 --- a/test/registered/observability/test_tracing_disaggregation.py +++ b/test/registered/observability/test_tracing_disaggregation.py @@ -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,