feat: add OpenTelemetry tracing to DiffGenerator (#21254)
This commit is contained in:
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user