From c6a6ba436812b330d2dd31e3d49314fb09c4cc4e Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Wed, 24 Dec 2025 19:51:26 +0800 Subject: [PATCH] Tiny refactor request logger (#15740) --- .../scheduler_runtime_checker_mixin.py | 7 +- .../managers/tokenizer_communicator_mixin.py | 38 ---- .../sglang/srt/managers/tokenizer_manager.py | 50 ++--- python/sglang/srt/utils/common.py | 49 ----- python/sglang/srt/utils/request_logger.py | 172 ++++++++++++++++++ .../debug_utils/test_request_logger.py | 60 ++++++ 6 files changed, 248 insertions(+), 128 deletions(-) create mode 100644 python/sglang/srt/utils/request_logger.py create mode 100644 test/registered/debug_utils/test_request_logger.py diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index 718f25d8b..60360574d 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -10,11 +10,8 @@ from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache -from sglang.srt.utils.common import ( - ceil_align, - disable_request_logging, - raise_error_or_warn, -) +from sglang.srt.utils.common import ceil_align, raise_error_or_warn +from sglang.srt.utils.request_logger import disable_request_logging from sglang.srt.utils.watchdog import WatchdogRaw if TYPE_CHECKING: diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index a5626f0b7..e5d42bed8 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -737,44 +737,6 @@ class TokenizerCommunicatorMixin: ): await self.send_to_scheduler.send_pyobj(obj) - def get_log_request_metadata(self): - max_length = None - skip_names = None - out_skip_names = None - if self.log_requests: - if self.log_requests_level == 0: - max_length = 1 << 30 - skip_names = { - "text", - "input_ids", - "input_embeds", - "image_data", - "audio_data", - "lora_path", - "sampling_params", - } - out_skip_names = {"text", "output_ids", "embedding"} - elif self.log_requests_level == 1: - max_length = 1 << 30 - skip_names = { - "text", - "input_ids", - "input_embeds", - "image_data", - "audio_data", - "lora_path", - } - out_skip_names = {"text", "output_ids", "embedding"} - elif self.log_requests_level == 2: - max_length = 2048 - elif self.log_requests_level == 3: - max_length = 1 << 30 - else: - raise ValueError( - f"Invalid --log-requests-level: {self.log_requests_level=}" - ) - return max_length, skip_names, out_skip_names - def _update_weight_version_if_provided(self, weight_version: Optional[str]) -> None: """Update weight version if provided.""" if weight_version is not None: diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 5fb310f2f..f6c5dfdc5 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -98,7 +98,6 @@ from sglang.srt.tracing.trace import ( ) from sglang.srt.utils import ( configure_gc_warning, - dataclass_to_string_truncated, freeze_gc, get_bool_env_var, get_or_create_event_loop, @@ -111,6 +110,7 @@ from sglang.srt.utils.hf_transformers_utils import ( get_tokenizer, get_tokenizer_from_processor, ) +from sglang.srt.utils.request_logger import RequestLogger from sglang.srt.utils.watchdog import Watchdog from sglang.utils import TypeBasedDispatcher, get_exception_traceback @@ -182,8 +182,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi # Parse args self.server_args = server_args self.enable_metrics = server_args.enable_metrics - self.log_requests = server_args.log_requests - self.log_requests_level = server_args.log_requests_level + self.request_logger = RequestLogger( + log_requests=server_args.log_requests, + log_requests_level=server_args.log_requests_level, + ) self.preferred_sampling_params = server_args.preferred_sampling_params self.crash_dump_folder = server_args.crash_dump_folder self.enable_trace = server_args.enable_trace @@ -308,12 +310,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi self.dump_requests_folder = "" # By default do not dump self.dump_requests_threshold = 1000 self.dump_request_list: List[Tuple] = [] - self.log_request_metadata = self.get_log_request_metadata() self.crash_dump_request_list: deque[Tuple] = deque() self.crash_dump_performed = False # Flag to ensure dump is only called once # Initialize performance metrics loggers with proper skip names - _, obj_skip_names, out_skip_names = self.log_request_metadata + _, obj_skip_names, out_skip_names = self.request_logger.metadata self.request_metrics_exporter_manager = RequestMetricsExporterManager( self.server_args, obj_skip_names, out_skip_names ) @@ -432,8 +433,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi self._trace_request_start(obj, created_time, request) if self.server_args.tokenizer_worker_num > 1: self._attach_multi_http_worker_info(obj) - if self.log_requests: - self._log_received_request(obj) + self.request_logger.log_received_request(obj, self.tokenizer) async with self.is_pause_cond: await self.is_pause_cond.wait_for(lambda: not self.is_pause) @@ -1044,13 +1044,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi out["meta_info"][ "response_sent_to_client_ts" ] = state.response_sent_to_client_ts - if self.log_requests: - max_length, skip_names, out_skip_names = self.log_request_metadata - if self.model_config.is_multimodal_gen: - msg = f"Finish: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}" - else: - msg = f"Finish: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}, out={dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}" - logger.info(msg) + self.request_logger.log_finished_request( + obj, out, is_multimodal_gen=self.model_config.is_multimodal_gen + ) if self.request_metrics_exporter_manager.exporter_enabled(): # Asynchronously write metrics for this request using the exporter manager. @@ -1315,10 +1311,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi return all_success, all_message, all_paused_requests def configure_logging(self, obj: ConfigureLoggingReq): - if obj.log_requests is not None: - self.log_requests = obj.log_requests - if obj.log_requests_level is not None: - self.log_requests_level = obj.log_requests_level + self.request_logger.configure( + log_requests=obj.log_requests, + log_requests_level=obj.log_requests_level, + ) if obj.dump_requests_folder is not None: self.dump_requests_folder = obj.dump_requests_folder if obj.dump_requests_threshold is not None: @@ -1326,7 +1322,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi if obj.crash_dump_folder is not None: self.crash_dump_folder = obj.crash_dump_folder logging.info(f"Config logging: {obj=}") - self.log_request_metadata = self.get_log_request_metadata() async def freeze_gc(self): """Send a freeze_gc message to the scheduler first, then freeze locally.""" @@ -2144,23 +2139,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi # Look up the LoRA ID from the registry and start tracking ongoing LoRA requests. obj.lora_id = await self.lora_registry.acquire(obj.lora_path) - def _log_received_request(self, obj: Union[GenerateReqInput, EmbeddingReqInput]): - max_length, skip_names, _ = self.log_request_metadata - logger.info( - f"Receive: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}" - ) - - # FIXME: This is a temporary fix to get the text from the input ids. - # We should remove this once we have a proper way. - if ( - self.log_requests_level >= 2 - and obj.text is None - and obj.input_ids is not None - and self.tokenizer is not None - ): - decoded = self.tokenizer.decode(obj.input_ids, skip_special_tokens=False) - obj.text = decoded - def _trace_request_start( self, obj: Union[GenerateReqInput, EmbeddingReqInput], diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index f8a2b4823..38f593ba9 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -18,7 +18,6 @@ import argparse import asyncio import builtins import ctypes -import dataclasses import functools import importlib import inspect @@ -66,7 +65,6 @@ from typing import ( Optional, Protocol, Sequence, - Set, Tuple, TypeVar, Union, @@ -2060,53 +2058,6 @@ def set_gpu_proc_affinity( logger.info(f"Process {pid} gpu_id {gpu_id} is running on CPUs: {p.cpu_affinity()}") -@lru_cache(maxsize=2) -def disable_request_logging() -> bool: - return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING") - - -def dataclass_to_string_truncated( - data, max_length=2048, skip_names: Optional[Set[str]] = None -): - if skip_names is None: - skip_names = set() - if isinstance(data, str): - if len(data) > max_length: - half_length = max_length // 2 - return f"{repr(data[:half_length])} ... {repr(data[-half_length:])}" - else: - return f"{repr(data)}" - elif isinstance(data, (list, tuple)): - if len(data) > max_length: - half_length = max_length // 2 - return str(data[:half_length]) + " ... " + str(data[-half_length:]) - else: - return str(data) - elif isinstance(data, dict): - return ( - "{" - + ", ".join( - f"'{k}': {dataclass_to_string_truncated(v, max_length)}" - for k, v in data.items() - if k not in skip_names - ) - + "}" - ) - elif dataclasses.is_dataclass(data): - fields = dataclasses.fields(data) - return ( - f"{data.__class__.__name__}(" - + ", ".join( - f"{f.name}={dataclass_to_string_truncated(getattr(data, f.name), max_length)}" - for f in fields - if f.name not in skip_names - ) - + ")" - ) - else: - return str(data) - - def permute_weight(x: torch.Tensor) -> torch.Tensor: b_ = x.shape[0] n_ = x.shape[1] diff --git a/python/sglang/srt/utils/request_logger.py b/python/sglang/srt/utils/request_logger.py new file mode 100644 index 000000000..7e2e13a04 --- /dev/null +++ b/python/sglang/srt/utils/request_logger.py @@ -0,0 +1,172 @@ +# Copyright 2023-2024 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 __future__ import annotations + +import dataclasses +import logging +from functools import lru_cache +from typing import TYPE_CHECKING, Any, Optional, Set, Tuple, Union + +from sglang.srt.utils.common import get_bool_env_var + +if TYPE_CHECKING: + from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput + +logger = logging.getLogger(__name__) + + +class RequestLogger: + def __init__(self, log_requests: bool, log_requests_level: int): + self.log_requests = log_requests + self.log_requests_level = log_requests_level + self.metadata: Tuple[Optional[int], Optional[Set[str]], Optional[Set[str]]] = ( + self._compute_metadata() + ) + + def configure( + self, + log_requests: Optional[bool] = None, + log_requests_level: Optional[int] = None, + ) -> None: + if log_requests is not None: + self.log_requests = log_requests + if log_requests_level is not None: + self.log_requests_level = log_requests_level + self.metadata = self._compute_metadata() + + def log_received_request( + self, obj: Union["GenerateReqInput", "EmbeddingReqInput"], tokenizer: Any = None + ) -> None: + if not self.log_requests: + return + + max_length, skip_names, _ = self.metadata + logger.info( + f"Receive: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}" + ) + + # FIXME: This is a temporary fix to get the text from the input ids. + # We should remove this once we have a proper way. + if ( + self.log_requests_level >= 2 + and obj.text is None + and obj.input_ids is not None + and tokenizer is not None + ): + decoded = tokenizer.decode(obj.input_ids, skip_special_tokens=False) + obj.text = decoded + + def log_finished_request( + self, + obj: Union["GenerateReqInput", "EmbeddingReqInput"], + out: Any, + is_multimodal_gen: bool = False, + ) -> None: + if not self.log_requests: + return + + max_length, skip_names, out_skip_names = self.metadata + if is_multimodal_gen: + msg = f"Finish: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}" + else: + msg = f"Finish: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}, out={_dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}" + logger.info(msg) + + def _compute_metadata( + self, + ) -> Tuple[Optional[int], Optional[Set[str]], Optional[Set[str]]]: + max_length: Optional[int] = None + skip_names: Optional[Set[str]] = None + out_skip_names: Optional[Set[str]] = None + if self.log_requests: + if self.log_requests_level == 0: + max_length = 1 << 30 + skip_names = { + "text", + "input_ids", + "input_embeds", + "image_data", + "audio_data", + "lora_path", + "sampling_params", + } + out_skip_names = {"text", "output_ids", "embedding"} + elif self.log_requests_level == 1: + max_length = 1 << 30 + skip_names = { + "text", + "input_ids", + "input_embeds", + "image_data", + "audio_data", + "lora_path", + } + out_skip_names = {"text", "output_ids", "embedding"} + elif self.log_requests_level == 2: + max_length = 2048 + elif self.log_requests_level == 3: + max_length = 1 << 30 + else: + raise ValueError( + f"Invalid --log-requests-level: {self.log_requests_level=}" + ) + return max_length, skip_names, out_skip_names + + +# TODO remove this? +@lru_cache(maxsize=2) +def disable_request_logging() -> bool: + return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING") + + +def _dataclass_to_string_truncated( + data: Any, max_length: int = 2048, skip_names: Optional[Set[str]] = None +) -> str: + if skip_names is None: + skip_names = set() + if isinstance(data, str): + if len(data) > max_length: + half_length = max_length // 2 + return f"{repr(data[:half_length])} ... {repr(data[-half_length:])}" + else: + return f"{repr(data)}" + elif isinstance(data, (list, tuple)): + if len(data) > max_length: + half_length = max_length // 2 + return str(data[:half_length]) + " ... " + str(data[-half_length:]) + else: + return str(data) + elif isinstance(data, dict): + return ( + "{" + + ", ".join( + f"'{k}': {_dataclass_to_string_truncated(v, max_length)}" + for k, v in data.items() + if k not in skip_names + ) + + "}" + ) + elif dataclasses.is_dataclass(data): + fields = dataclasses.fields(data) + return ( + f"{data.__class__.__name__}(" + + ", ".join( + f"{f.name}={_dataclass_to_string_truncated(getattr(data, f.name), max_length)}" + for f in fields + if f.name not in skip_names + ) + + ")" + ) + else: + return str(data) diff --git a/test/registered/debug_utils/test_request_logger.py b/test/registered/debug_utils/test_request_logger.py new file mode 100644 index 000000000..5b5de52c9 --- /dev/null +++ b/test/registered/debug_utils/test_request_logger.py @@ -0,0 +1,60 @@ +import io +import unittest + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=60, suite="nightly-1-gpu", nightly=True) + + +class TestRequestLogger(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.stdout = io.StringIO() + cls.stderr = io.StringIO() + + cls.process = popen_launch_server( + "Qwen/Qwen3-0.6B", + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--log-requests", + "--log-requests-level", + "2", + "--skip-server-warmup", + ], + return_stdout_stderr=(cls.stdout, cls.stderr), + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + cls.stdout.close() + cls.stderr.close() + + def test_request_logging(self): + response = requests.post( + DEFAULT_URL_FOR_TEST + "/generate", + json={ + "text": "Hello", + "sampling_params": {"max_new_tokens": 8, "temperature": 0}, + }, + timeout=30, + ) + self.assertEqual(response.status_code, 200) + + combined_output = self.stdout.getvalue() + self.stderr.getvalue() + self.assertIn("Receive:", combined_output) + self.assertIn("Finish:", combined_output) + + +if __name__ == "__main__": + unittest.main()