Tiny extract file logging utils (#16870)
This commit is contained in:
@@ -0,0 +1,72 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import socket
|
||||||
|
import sys
|
||||||
|
from datetime import datetime
|
||||||
|
from logging.handlers import TimedRotatingFileHandler
|
||||||
|
from typing import List, Optional, Union
|
||||||
|
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
|
||||||
|
def create_log_targets(
|
||||||
|
*, targets: Optional[List[str]], name_prefix: str
|
||||||
|
) -> List[logging.Logger]:
|
||||||
|
if not targets:
|
||||||
|
return [_create_log_target_stdout(name_prefix)]
|
||||||
|
return [_create_log_target(t, name_prefix) for t in targets]
|
||||||
|
|
||||||
|
|
||||||
|
def _create_log_target(target: str, name_prefix: str) -> logging.Logger:
|
||||||
|
if target.lower() == "stdout":
|
||||||
|
return _create_log_target_stdout(name_prefix)
|
||||||
|
return _create_log_target_file(target, name_prefix)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_log_target_stdout(name_prefix: str) -> logging.Logger:
|
||||||
|
return _create_logger_with_handler(
|
||||||
|
f"{name_prefix}.stdout", logging.StreamHandler(sys.stdout)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_log_target_file(directory: str, name_prefix: str) -> logging.Logger:
|
||||||
|
os.makedirs(directory, exist_ok=True)
|
||||||
|
hostname = socket.gethostname()
|
||||||
|
rank = dist.get_rank() if dist.is_initialized() else 0
|
||||||
|
filename = os.path.join(directory, f"{hostname}_{rank}.log")
|
||||||
|
handler = TimedRotatingFileHandler(
|
||||||
|
filename, when="H", backupCount=0, encoding="utf-8"
|
||||||
|
)
|
||||||
|
return _create_logger_with_handler(
|
||||||
|
f"{name_prefix}.file.{directory}.{hostname}_{rank}", handler
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_logger_with_handler(name: str, handler: logging.Handler) -> logging.Logger:
|
||||||
|
logger = logging.getLogger(name)
|
||||||
|
logger.setLevel(logging.INFO)
|
||||||
|
logger.propagate = False
|
||||||
|
if not logger.handlers:
|
||||||
|
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||||
|
logger.addHandler(handler)
|
||||||
|
return logger
|
||||||
|
|
||||||
|
|
||||||
|
def log_json(
|
||||||
|
loggers: Union[logging.Logger, List[logging.Logger]], event: str, data: dict
|
||||||
|
) -> None:
|
||||||
|
log_data = {
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"event": event,
|
||||||
|
**data,
|
||||||
|
}
|
||||||
|
msg = json.dumps(log_data, ensure_ascii=False)
|
||||||
|
|
||||||
|
if not isinstance(loggers, list):
|
||||||
|
loggers = [loggers]
|
||||||
|
|
||||||
|
for logger in loggers:
|
||||||
|
logger.info(msg)
|
||||||
@@ -14,19 +14,13 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
import socket
|
|
||||||
from datetime import datetime
|
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from logging.handlers import TimedRotatingFileHandler
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union
|
||||||
|
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.utils.common import get_bool_env_var
|
from sglang.srt.utils.common import get_bool_env_var
|
||||||
|
from sglang.srt.utils.log_utils import create_log_targets, log_json
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
import fastapi
|
import fastapi
|
||||||
@@ -67,9 +61,9 @@ class RequestLogger:
|
|||||||
self.log_exceeded_ms = envs.SGLANG_LOG_REQUEST_EXCEEDED_MS.get()
|
self.log_exceeded_ms = envs.SGLANG_LOG_REQUEST_EXCEEDED_MS.get()
|
||||||
|
|
||||||
def _setup_targets(self) -> List[logging.Logger]:
|
def _setup_targets(self) -> List[logging.Logger]:
|
||||||
if not self.log_requests_target:
|
return create_log_targets(
|
||||||
return [_create_log_target_stdout()]
|
targets=self.log_requests_target, name_prefix=__name__
|
||||||
return [_create_log_target(t) for t in self.log_requests_target]
|
)
|
||||||
|
|
||||||
def configure(
|
def configure(
|
||||||
self,
|
self,
|
||||||
@@ -108,7 +102,7 @@ class RequestLogger:
|
|||||||
}
|
}
|
||||||
if headers:
|
if headers:
|
||||||
log_data["headers"] = headers
|
log_data["headers"] = headers
|
||||||
self._log_json("request.received", log_data)
|
log_json(self.targets, "request.received", log_data)
|
||||||
else:
|
else:
|
||||||
headers_str = f", headers={headers}" if headers else ""
|
headers_str = f", headers={headers}" if headers else ""
|
||||||
self._log(
|
self._log(
|
||||||
@@ -153,7 +147,7 @@ class RequestLogger:
|
|||||||
log_data["out"] = _transform_data_for_logging(
|
log_data["out"] = _transform_data_for_logging(
|
||||||
out, max_length, out_skip_names
|
out, max_length, out_skip_names
|
||||||
)
|
)
|
||||||
self._log_json("request.finished", log_data)
|
log_json(self.targets, "request.finished", log_data)
|
||||||
else:
|
else:
|
||||||
obj_str = _dataclass_to_string_truncated(
|
obj_str = _dataclass_to_string_truncated(
|
||||||
obj, max_length, skip_names=skip_names
|
obj, max_length, skip_names=skip_names
|
||||||
@@ -206,14 +200,6 @@ class RequestLogger:
|
|||||||
)
|
)
|
||||||
return max_length, skip_names, out_skip_names
|
return max_length, skip_names, out_skip_names
|
||||||
|
|
||||||
def _log_json(self, event: str, data: dict) -> None:
|
|
||||||
log_data = {
|
|
||||||
"timestamp": datetime.now().isoformat(),
|
|
||||||
"event": event,
|
|
||||||
**data,
|
|
||||||
}
|
|
||||||
self._log(json.dumps(log_data, ensure_ascii=False))
|
|
||||||
|
|
||||||
def _log(self, msg: str) -> None:
|
def _log(self, msg: str) -> None:
|
||||||
for target in self.targets:
|
for target in self.targets:
|
||||||
target.info(msg)
|
target.info(msg)
|
||||||
@@ -225,39 +211,6 @@ def disable_request_logging() -> bool:
|
|||||||
return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING")
|
return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING")
|
||||||
|
|
||||||
|
|
||||||
def _create_log_target(target: str) -> logging.Logger:
|
|
||||||
if target.lower() == "stdout":
|
|
||||||
return _create_log_target_stdout()
|
|
||||||
return _create_log_target_file(target)
|
|
||||||
|
|
||||||
|
|
||||||
def _create_log_target_stdout() -> logging.Logger:
|
|
||||||
return _create_logger_with_handler(f"{__name__}.stdout", logging.StreamHandler())
|
|
||||||
|
|
||||||
|
|
||||||
def _create_log_target_file(directory: str) -> logging.Logger:
|
|
||||||
os.makedirs(directory, exist_ok=True)
|
|
||||||
hostname = socket.gethostname()
|
|
||||||
rank = dist.get_rank() if dist.is_initialized() else 0
|
|
||||||
filename = os.path.join(directory, f"{hostname}_{rank}.log")
|
|
||||||
handler = TimedRotatingFileHandler(
|
|
||||||
filename, when="H", backupCount=0, encoding="utf-8"
|
|
||||||
)
|
|
||||||
return _create_logger_with_handler(
|
|
||||||
f"{__name__}.file.{directory}.{hostname}_{rank}", handler
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _create_logger_with_handler(name: str, handler: logging.Handler) -> logging.Logger:
|
|
||||||
logger = logging.getLogger(name)
|
|
||||||
logger.setLevel(logging.INFO)
|
|
||||||
logger.propagate = False
|
|
||||||
if not logger.handlers:
|
|
||||||
handler.setFormatter(logging.Formatter("%(message)s"))
|
|
||||||
logger.addHandler(handler)
|
|
||||||
return logger
|
|
||||||
|
|
||||||
|
|
||||||
# TODO unify this w/ `_transform_data_for_logging` if we find performance enough
|
# TODO unify this w/ `_transform_data_for_logging` if we find performance enough
|
||||||
def _dataclass_to_string_truncated(
|
def _dataclass_to_string_truncated(
|
||||||
data: Any, max_length: int = 2048, skip_names: Optional[Set[str]] = None
|
data: Any, max_length: int = 2048, skip_names: Optional[Set[str]] = None
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
import io
|
||||||
|
import json
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
import uuid
|
||||||
|
from contextlib import redirect_stdout
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from sglang.srt.utils.log_utils import create_log_targets, log_json
|
||||||
|
|
||||||
|
|
||||||
|
class TestLogUtils(unittest.TestCase):
|
||||||
|
def test_stdout(self):
|
||||||
|
for targets in [["stdout"], None]:
|
||||||
|
with self.subTest(targets=targets):
|
||||||
|
buf = io.StringIO()
|
||||||
|
with redirect_stdout(buf):
|
||||||
|
loggers = create_log_targets(
|
||||||
|
targets=targets, name_prefix=f"test_stdout_{uuid.uuid4()}"
|
||||||
|
)
|
||||||
|
self.assertEqual(len(loggers), 1)
|
||||||
|
log_json(loggers[0], "test.event", {"key": "value"})
|
||||||
|
data = json.loads(buf.getvalue().strip())
|
||||||
|
self.assertIn("timestamp", data)
|
||||||
|
self.assertEqual(data["event"], "test.event")
|
||||||
|
self.assertEqual(data["key"], "value")
|
||||||
|
|
||||||
|
def test_file(self):
|
||||||
|
with tempfile.TemporaryDirectory() as temp_dir:
|
||||||
|
loggers = create_log_targets(
|
||||||
|
targets=[temp_dir], name_prefix=f"test_file_{uuid.uuid4()}"
|
||||||
|
)
|
||||||
|
self.assertEqual(len(loggers), 1)
|
||||||
|
log_json(loggers, "file.event", {"data": 123})
|
||||||
|
_flush_all(loggers)
|
||||||
|
data = _read_log_file(temp_dir)
|
||||||
|
self.assertIn("timestamp", data)
|
||||||
|
self.assertEqual(data["event"], "file.event")
|
||||||
|
self.assertEqual(data["data"], 123)
|
||||||
|
|
||||||
|
def test_multiple_targets(self):
|
||||||
|
with tempfile.TemporaryDirectory() as temp_dir:
|
||||||
|
buf = io.StringIO()
|
||||||
|
with redirect_stdout(buf):
|
||||||
|
loggers = create_log_targets(
|
||||||
|
targets=["stdout", temp_dir],
|
||||||
|
name_prefix=f"test_multi_{uuid.uuid4()}",
|
||||||
|
)
|
||||||
|
self.assertEqual(len(loggers), 2)
|
||||||
|
log_json(loggers, "multi.event", {"x": 1})
|
||||||
|
_flush_all(loggers)
|
||||||
|
stdout_data = json.loads(buf.getvalue().strip())
|
||||||
|
file_data = _read_log_file(temp_dir)
|
||||||
|
self.assertEqual(stdout_data["event"], "multi.event")
|
||||||
|
self.assertEqual(file_data["event"], "multi.event")
|
||||||
|
self.assertEqual(stdout_data["x"], file_data["x"])
|
||||||
|
|
||||||
|
|
||||||
|
def _flush_all(loggers: list) -> None:
|
||||||
|
for logger in loggers:
|
||||||
|
for handler in logger.handlers:
|
||||||
|
handler.flush()
|
||||||
|
|
||||||
|
|
||||||
|
def _read_log_file(temp_dir: str) -> dict:
|
||||||
|
log_files = list(Path(temp_dir).glob("*.log"))
|
||||||
|
assert len(log_files) == 1
|
||||||
|
return json.loads(log_files[0].read_text().strip())
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user