diff --git a/docs_new/docs/sglang-diffusion/api/cli.mdx b/docs_new/docs/sglang-diffusion/api/cli.mdx index ae1229f63..f88f378e3 100644 --- a/docs_new/docs/sglang-diffusion/api/cli.mdx +++ b/docs_new/docs/sglang-diffusion/api/cli.mdx @@ -110,6 +110,13 @@ For quantized transformer checkpoints, prefer: See [Quantization](../quantization) for supported quantization families and examples. +### Request logging + +- `--log-requests`: Log user-facing fields of all requests (default: `False`). The verbosity is decided by `--log-requests-level`. +- `--log-requests-level {0|1|2|3}`: Verbosity level for request logging (default: `2`). 0: Log metadata (request id). 1: Log metadata and sampling config (seed, steps, guidance, resolution, frames, fps, ...). 2: Log metadata, sampling config and prompt (truncated to 2 KiB). 3: Log metadata, sampling config and full prompt. +- `--log-requests-format {text|json}`: Format for request logging (default: `text`). `text` is human-readable; `json` outputs structured JSON lines. +- `--log-requests-target {TARGET...}`: Target(s) for request logging. Use `stdout` for console output and/or directory path(s) for file output. Can specify multiple targets, e.g., `--log-requests-target stdout /my/log/dir`. + ## Configuration Files Use `--config` to load JSON or YAML configuration. Command-line flags override values from the config file. diff --git a/python/sglang/multimodal_gen/runtime/scheduler_client.py b/python/sglang/multimodal_gen/runtime/scheduler_client.py index 2455d5767..ef78bd592 100644 --- a/python/sglang/multimodal_gen/runtime/scheduler_client.py +++ b/python/sglang/multimodal_gen/runtime/scheduler_client.py @@ -1,6 +1,6 @@ import pickle import time -from typing import Any +from typing import Any, Optional import zmq import zmq.asyncio @@ -9,6 +9,9 @@ from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.request_logger import ( + DiffusionRequestLogger, +) logger = init_logger(__name__) @@ -56,6 +59,7 @@ class SchedulerClient: self.context = None self.scheduler_socket = None self.server_args = None + self.request_logger: Optional[DiffusionRequestLogger] = None def initialize(self, server_args: ServerArgs): if self.context is not None and not self.context.closed: @@ -63,6 +67,7 @@ class SchedulerClient: self.close() self.server_args = server_args + self.request_logger = DiffusionRequestLogger.from_server_args(server_args) self.context = zmq.Context() self.scheduler_socket = self.context.socket(zmq.REQ) @@ -80,6 +85,7 @@ class SchedulerClient: def forward(self, batch: Any, timeout_ms: int | None = None) -> Any: """Sends a batch or request to the scheduler and waits for the response.""" + self.request_logger.log_received_request(batch) previous_timeout_ms = None if timeout_ms is not None: previous_timeout_ms = self.scheduler_socket.getsockopt(zmq.RCVTIMEO) @@ -88,6 +94,7 @@ class SchedulerClient: self.scheduler_socket.send_pyobj(batch) output_batch = self.scheduler_socket.recv_pyobj() _materialize_output_batch_file_refs(output_batch) + self.request_logger.log_finished_request(batch, output_batch) return output_batch except zmq.error.Again: logger.error("Timeout waiting for response from scheduler.") @@ -142,6 +149,7 @@ class AsyncSchedulerClient: def __init__(self): self.context = None self.server_args = None + self.request_logger: Optional[DiffusionRequestLogger] = None def initialize(self, server_args: ServerArgs): if self.context is not None and not self.context.closed: @@ -151,11 +159,13 @@ class AsyncSchedulerClient: self.close() self.server_args = server_args + self.request_logger = DiffusionRequestLogger.from_server_args(server_args) self.context = zmq.asyncio.Context() logger.debug("AsyncSchedulerClient initialized with zmq.asyncio.Context") async def forward(self, batch: Any) -> Any: """Sends a batch or request to the scheduler and waits for the response.""" + self.request_logger.log_received_request(batch) if self.context is None: raise RuntimeError( "AsyncSchedulerClient is not initialized. Call initialize() first." @@ -175,6 +185,7 @@ class AsyncSchedulerClient: payload = await socket.recv() output_batch = pickle.loads(payload) _materialize_output_batch_file_refs(output_batch) + self.request_logger.log_finished_request(batch, output_batch) return output_batch except zmq.error.Again: logger.error("Timeout waiting for response from scheduler.") diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index d8d3dde2c..2c899b8b8 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -14,7 +14,7 @@ import sys import tempfile from dataclasses import field from enum import Enum -from typing import Any, Literal, Optional +from typing import Any, List, Literal, Optional import addict import yaml @@ -343,6 +343,10 @@ class ServerArgs(DisaggServerArgsMixin): # Logging log_level: str = "info" + log_requests: bool = False + log_requests_level: int = 2 + log_requests_format: str = "text" + log_requests_target: Optional[List[str]] = None uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list) # Tracing @@ -1652,6 +1656,38 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.otlp_traces_endpoint, help="OTLP collector endpoint when --enable-trace is set. Format: :", ) + parser.add_argument( + "--log-requests", + action="store_true", + help="Log user-facing fields of all requests (default: False). " + "Verbosity is controlled by --log-requests-level.", + ) + parser.add_argument( + "--log-requests-level", + type=int, + default=ServerArgs.log_requests_level, + choices=[0, 1, 2, 3], + help="Verbosity level for request logging. " + "0: Log request metadata only (request_id). " + "1: Log metadata + sampling config (seed, steps, guidance, resolution, frames, fps, ...). " + "2: Log metadata + sampling config + prompt/negative prompt (truncated to 2 KiB). " + "3: Log metadata + sampling config + full prompt/negative prompt.", + ) + parser.add_argument( + "--log-requests-format", + type=str, + default=ServerArgs.log_requests_format, + choices=["text", "json"], + help="Format for request logging: 'text' (human-readable) or 'json' (structured)", + ) + parser.add_argument( + "--log-requests-target", + type=str, + nargs="+", + default=ServerArgs.log_requests_target, + help="Target(s) for request logging: 'stdout' and/or directory path(s) for file output. " + "Can specify multiple targets, e.g., '--log-requests-target stdout /my/path'. ", + ) parser.add_argument( "--uvicorn-access-log-exclude-prefixes", type=str, diff --git a/python/sglang/multimodal_gen/runtime/utils/request_logger.py b/python/sglang/multimodal_gen/runtime/utils/request_logger.py new file mode 100644 index 000000000..1b2f20cc4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/request_logger.py @@ -0,0 +1,204 @@ +# Copyright 2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + + +from typing import Any, Optional + +from sglang.srt.environ import envs +from sglang.srt.utils.log_utils import create_log_targets, log_json +from sglang.srt.utils.request_logger import ( + _dataclass_to_string_truncated, + _transform_data_for_logging, +) + +# Core generation knobs logged per record. Prompt text is logged separately, +# gated by the level, so it is excluded here. +_SAMPLING_CONFIG_FIELDS = ( + "data_type", + "seed", + "num_inference_steps", + "guidance_scale", + "true_cfg_scale", + "width", + "height", + "num_frames", + "fps", + "num_outputs_per_prompt", +) + +# Level 2 truncates prompt text; level 3 keeps it whole. Lower levels log no +# prompt at all. +_TRUNCATE_LENGTH = 2048 +_UNLIMITED = 1 << 30 + + +class DiffusionRequestLogger: + def __init__( + self, + log_requests: bool, + log_requests_level: int, + log_requests_format: str, + log_requests_target: Optional[list], + ): + self.log_requests = log_requests + self.log_requests_level = log_requests_level + self.log_requests_format = log_requests_format + self.log_requests_target = log_requests_target + self.targets = create_log_targets( + targets=log_requests_target, name_prefix=__name__ + ) + self.log_exceeded_ms = envs.SGLANG_LOG_REQUEST_EXCEEDED_MS.get() + self._max_length = ( + _TRUNCATE_LENGTH if self.log_requests_level == 2 else _UNLIMITED + ) + + @classmethod + def from_server_args(cls, server_args: Any) -> "DiffusionRequestLogger": + """Build a logger from server args.""" + return cls( + log_requests=server_args.log_requests, + log_requests_level=server_args.log_requests_level, + log_requests_format=server_args.log_requests_format, + log_requests_target=server_args.log_requests_target, + ) + + @staticmethod + def _request_id(req: Any) -> Optional[str]: + """The request's id, or ``None`` if absent.""" + return getattr(req, "request_id", None) + + def _config_view(self, req: Any, *, drop_seed: bool = False) -> dict: + """Sampling config + (level >= 2) prompt, gated by the log level. + Returns ``{}`` below level 1.""" + sp = getattr(req, "sampling_params", None) + if sp is None or self.log_requests_level < 1: + return {} + cfg = {name: getattr(sp, name, None) for name in _SAMPLING_CONFIG_FIELDS} + if drop_seed: + cfg.pop("seed", None) + view: dict = {"sampling_params": cfg} + if self.log_requests_level >= 2: + view["prompt"] = getattr(sp, "prompt", None) + view["negative_prompt"] = getattr(sp, "negative_prompt", None) + return view + + def _result_view(self, result: Any) -> dict: + """Result-side fields for a finished record: latency and error.""" + e2e_latency = 0.0 + metrics = getattr(result, "metrics", None) if result is not None else None + if metrics is not None: + e2e_latency = getattr(metrics, "total_duration_s", 0.0) or 0.0 + return { + "meta_info": {"e2e_latency": e2e_latency}, + "error": getattr(result, "error", None) if result is not None else None, + } + + def _emit(self, msg: str) -> None: + for target in self.targets: + target.info(msg) + + def _per_request_view(self, req: Any) -> dict: + """Per-output identity within a batch: ``request_id`` plus ``seed`` at + level >= 1.""" + sp = getattr(req, "sampling_params", None) + view: dict = {"request_id": self._request_id(req)} + if self.log_requests_level >= 1: + view["seed"] = getattr(sp, "seed", None) if sp is not None else None + return view + + def _batch_record(self, reqs: list) -> tuple: + """Build the ``rid`` / ``obj`` for one forward call""" + rids = [self._request_id(req) or "unknown" for req in reqs] + if len(reqs) == 1: + # Single request: scalar rid + flat dict obj (id + config). + req = reqs[0] + obj = {"request_id": self._request_id(req), **self._config_view(req)} + return rids[0], obj + + shared_views = [self._config_view(req, drop_seed=True) for req in reqs] + if all(view == shared_views[0] for view in shared_views): + obj = { + **shared_views[0], + "outputs": [self._per_request_view(req) for req in reqs], + } + else: + # Configs genuinely differ: list each request's full payload verbatim. + obj = [ + {"request_id": self._request_id(req), **self._config_view(req)} + for req in reqs + ] + return rids, obj + + def _loggable(self, req: Any) -> bool: + """Whether ``req`` should be recorded: logging is on, it's a real + generation request (control messages -- LoRA / weight / stats / shutdown + -- have no ``sampling_params`` and are skipped), and it's not a warmup.""" + return ( + self.log_requests + and getattr(req, "sampling_params", None) is not None + and not getattr(req, "is_warmup", False) + ) + + def _logged_reqs(self, batch: Any) -> list: + """Normalize ``batch`` to a list and drop control / warmup requests.""" + reqs = batch if isinstance(batch, (list, tuple)) else [batch] + return [r for r in reqs if self._loggable(r)] + + def log_received_request(self, batch: Any) -> None: + reqs = self._logged_reqs(batch) + if not reqs: + return + + rid, obj = self._batch_record(reqs) + max_length = self._max_length + + if self.log_requests_format == "json": + log_json( + self.targets, + "request.received", + {"rid": rid, "obj": _transform_data_for_logging(obj, max_length)}, + ) + else: + self._emit( + f"Receive: obj={_dataclass_to_string_truncated(obj, max_length)}" + ) + + def log_finished_request(self, batch: Any, result: Any) -> None: + reqs = self._logged_reqs(batch) + if not reqs: + return + + out = self._result_view(result) + e2e_latency_ms = out["meta_info"]["e2e_latency"] * 1000 + if self.log_exceeded_ms > 0 and e2e_latency_ms < self.log_exceeded_ms: + return + + rid, obj = self._batch_record(reqs) + max_length = self._max_length + + if self.log_requests_format == "json": + log_json( + self.targets, + "request.finished", + { + "rid": rid, + "obj": _transform_data_for_logging(obj, max_length), + "out": _transform_data_for_logging(out, max_length), + }, + ) + else: + self._emit( + f"Finish: obj={_dataclass_to_string_truncated(obj, max_length)}" + f", out={_dataclass_to_string_truncated(out, max_length)}" + ) diff --git a/python/sglang/multimodal_gen/test/server/test_request_logger.py b/python/sglang/multimodal_gen/test/server/test_request_logger.py new file mode 100644 index 000000000..404ee0c5b --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/test_request_logger.py @@ -0,0 +1,233 @@ +""" +Test request logging for diffusion models. + +Tests the --log-requests CLI flags for diffusion model serving, +verifying that request logs are correctly written to stdout and files. +""" + +import json +import os +import shutil +import tempfile +import time +from pathlib import Path + +import pytest +from openai import OpenAI + +from sglang.multimodal_gen.test.server.test_server_utils import ServerManager +from sglang.multimodal_gen.test.test_utils import get_dynamic_server_port + +# Test models and prompts +IMAGE_MODEL = "Efficient-Large-Model/Sana_600M_512px_diffusers" +IMAGE_PROMPT = "A beautiful sunset over mountains, oil painting style" +VIDEO_MODEL = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" +VIDEO_PROMPT = "A cat playing with a ball" + +# Timeout settings +BASE_TIMEOUT = float(os.environ.get("SGLANG_TEST_OPENAI_REQUEST_TIMEOUT_SECS", "600")) +POLL_INTERVAL = 1.0 + + +def _start_server(model: str, log_format: str): + """Start server with request logging enabled.""" + temp_dir = tempfile.mkdtemp() + port = get_dynamic_server_port() + extra_args = ( + f"--log-requests " + f"--log-requests-level 2 " + f"--log-requests-format {log_format} " + f"--log-requests-target stdout {temp_dir} " + f"--strict-ports" + ) + wait_deadline = float(os.environ.get("SGLANG_TEST_WAIT_SECS", "1200")) + manager = ServerManager( + model=model, + port=port, + wait_deadline=wait_deadline, + extra_args=extra_args, + ) + ctx = manager.start() + ctx.temp_dir = temp_dir + return ctx + + +def _cleanup_server(ctx): + """Cleanup server and temp directory.""" + ctx.cleanup() + shutil.rmtree(ctx.temp_dir, ignore_errors=True) + + +def _create_client(ctx) -> OpenAI: + """Create OpenAI client for the server.""" + return OpenAI( + api_key="test", + base_url=f"http://localhost:{ctx.port}/v1", + timeout=BASE_TIMEOUT, + ) + + +def _wait_for_video_completion(client: OpenAI, video_id: str, timeout: float): + """Poll video job until completion.""" + deadline = time.time() + timeout + while time.time() < deadline: + page = client.videos.list() + item = next((v for v in page.data if v.id == video_id), None) + status = getattr(item, "status", None) if item else None + + if status == "completed": + return True + if status in ("failed", "cancelled", "deleted"): + pytest.fail(f"Video job {video_id} ended with status={status}") + + time.sleep(POLL_INTERVAL) + + pytest.fail(f"Video job {video_id} did not complete in {timeout}s") + + +def _verify_json_logs(content: str): + """Verify JSON logs contain request.received and request.finished events.""" + has_received = False + has_finished = False + + for line in content.splitlines(): + idx = line.find("{") + if idx == -1: + continue + try: + data = json.loads(line[idx:]) + except json.JSONDecodeError: + continue + + if data.get("event") == "request.received": + has_received = True + elif data.get("event") == "request.finished": + has_finished = True + + assert has_received, "request.received event not found" + assert has_finished, "request.finished event not found" + + +def _verify_text_logs(content: str, prompt: str): + """Verify text logs contain Receive, prompt, and Finish markers.""" + assert "Receive:" in content, "'Receive:' not found" + assert prompt in content, f"Prompt '{prompt}' not found" + assert "Finish:" in content, "'Finish:' not found" + + +@pytest.fixture(scope="class") +def image_text_server(): + """Server with text-format logging for image model.""" + ctx = _start_server(IMAGE_MODEL, "text") + yield ctx + _cleanup_server(ctx) + + +@pytest.fixture(scope="class") +def image_json_server(): + """Server with JSON-format logging for image model.""" + ctx = _start_server(IMAGE_MODEL, "json") + yield ctx + _cleanup_server(ctx) + + +@pytest.fixture(scope="class") +def video_text_server(): + """Server with text-format logging for video model.""" + ctx = _start_server(VIDEO_MODEL, "text") + yield ctx + _cleanup_server(ctx) + + +@pytest.fixture(scope="class") +def video_json_server(): + """Server with JSON-format logging for video model.""" + ctx = _start_server(VIDEO_MODEL, "json") + yield ctx + _cleanup_server(ctx) + + +class TestImageRequestLoggerText: + """Test text-format request logging for image models.""" + + def test_request_logging(self, image_text_server): + ctx = image_text_server + client = _create_client(ctx) + + # Image generation is synchronous, waits for completion + client.images.generate(prompt=IMAGE_PROMPT, size="256x256", n=1) + + # Verify stdout and file logs + stdout = ctx.log_tail(lines=500) + _verify_text_logs(stdout, IMAGE_PROMPT[:30]) + + logs = list(Path(ctx.temp_dir).glob("*.log")) + assert logs, "No log files found" + _verify_text_logs("".join(f.read_text() for f in logs), IMAGE_PROMPT[:30]) + + +class TestImageRequestLoggerJson: + """Test JSON-format request logging for image models.""" + + def test_request_logging(self, image_json_server): + ctx = image_json_server + client = _create_client(ctx) + + # Image generation is synchronous, waits for completion + client.images.generate(prompt=IMAGE_PROMPT, size="256x256", n=1) + + # Verify stdout and file logs + stdout = ctx.log_tail(lines=500) + _verify_json_logs(stdout) + + logs = list(Path(ctx.temp_dir).glob("*.log")) + assert logs, "No log files found" + _verify_json_logs("".join(f.read_text() for f in logs)) + + +class TestVideoRequestLoggerText: + """Test text-format request logging for video models.""" + + def test_request_logging(self, video_text_server): + ctx = video_text_server + client = _create_client(ctx) + + # Video generation is async - create job and poll until completion + job = client.videos.create( + prompt=VIDEO_PROMPT, + size="832x480", + extra_body={"num_frames": 5, "num_inference_steps": 10}, + ) + _wait_for_video_completion(client, job.id, BASE_TIMEOUT * 2) + + # Verify stdout and file logs + stdout = ctx.log_tail(lines=500) + _verify_text_logs(stdout, VIDEO_PROMPT[:20]) + + logs = list(Path(ctx.temp_dir).glob("*.log")) + assert logs, "No log files found" + _verify_text_logs("".join(f.read_text() for f in logs), VIDEO_PROMPT[:20]) + + +class TestVideoRequestLoggerJson: + """Test JSON-format request logging for video models.""" + + def test_request_logging(self, video_json_server): + ctx = video_json_server + client = _create_client(ctx) + + # Video generation is async - create job and poll until completion + job = client.videos.create( + prompt=VIDEO_PROMPT, + size="832x480", + extra_body={"num_frames": 5, "num_inference_steps": 10}, + ) + _wait_for_video_completion(client, job.id, BASE_TIMEOUT * 2) + + # Verify stdout and file logs + stdout = ctx.log_tail(lines=500) + _verify_json_logs(stdout) + + logs = list(Path(ctx.temp_dir).glob("*.log")) + assert logs, "No log files found" + _verify_json_logs("".join(f.read_text() for f in logs))