|
1 | 1 | import logging |
| 2 | +import numbers |
2 | 3 | import sys |
3 | 4 | from copy import copy |
| 5 | +from typing import Any |
4 | 6 |
|
| 7 | +import orjson |
5 | 8 | from icij_common.logging_utils import DATE_FMT, STREAM_HANDLER_FMT |
6 | | -from pythonjsonlogger.core import RESERVED_ATTRS, BaseJsonFormatter |
| 9 | +from pythonjsonlogger.core import BaseJsonFormatter |
7 | 10 | from pythonjsonlogger.orjson import OrjsonFormatter |
8 | 11 | from temporalio import activity, workflow |
9 | 12 |
|
10 | | -from .config import LogLevel |
| 13 | +from .config import LogFormat, LogLevel |
11 | 14 | from .interceptors import get_trace_context |
12 | 15 |
|
| 16 | +_BASE_ATTRS = [ |
| 17 | + "asctime", |
| 18 | + "exc_info", |
| 19 | + "filename", |
| 20 | + "funcName", |
| 21 | + "levelname", |
| 22 | + "levelno", |
| 23 | + "lineno", |
| 24 | + "module", |
| 25 | + "msecs", |
| 26 | + "message", |
| 27 | + "msg", |
| 28 | + "name", |
| 29 | + "pathname", |
| 30 | +] |
13 | 31 | _ACT_LOGGER_ATTRS = ["activity_type", "activity_id", "activity_run_id"] |
14 | 32 | _WF_LOGGED_ATTRS = ["workflow_type", "workflow_id", "workflow_run_id"] |
15 | 33 | _TRACE_CONTEXT_ATTRS = ["trace_id", "parent_id", "traceparent"] |
| 34 | + |
16 | 35 | _LOGGED_ATTRIBUTES = ( |
17 | | - copy(RESERVED_ATTRS) |
| 36 | + copy(_BASE_ATTRS) |
18 | 37 | + _WF_LOGGED_ATTRS |
19 | 38 | + _ACT_LOGGER_ATTRS |
20 | 39 | + _TRACE_CONTEXT_ATTRS |
|
28 | 47 |
|
29 | 48 |
|
30 | 49 | def setup_worker_loggers( |
31 | | - loggers: dict[str, LogLevel], *, worker_id: str | None, in_json: bool |
| 50 | + loggers: dict[str, LogLevel], *, worker_id: str | None, format: LogFormat |
32 | 51 | ) -> None: |
33 | 52 | worker_filter = WorkerFilter(worker_id) |
34 | 53 | for logger_name, level_str in loggers.items(): |
35 | 54 | level = getattr(logging, level_str) |
36 | 55 | logger = logging.getLogger(logger_name) |
37 | 56 | logger.setLevel(level) |
38 | 57 | logger.handlers = [] |
39 | | - for handler in _get_worker_handlers(level, worker_filter, in_json=in_json): |
| 58 | + for handler in _get_worker_handlers(level, worker_filter, format=format): |
40 | 59 | logger.addHandler(handler) |
41 | 60 |
|
42 | 61 |
|
@@ -64,23 +83,50 @@ def filter(self, record: logging.LogRecord) -> bool: |
64 | 83 |
|
65 | 84 |
|
66 | 85 | def _get_worker_handlers( |
67 | | - level: int, worker_filter: WorkerFilter, *, in_json: bool |
| 86 | + level: int, worker_filter: WorkerFilter, *, format: LogFormat |
68 | 87 | ) -> list[logging.Handler]: |
69 | 88 | stream_handler = logging.StreamHandler(sys.stderr) |
70 | | - if in_json: |
71 | | - fmt = _json_formatter(datefmt=DATE_FMT) |
72 | | - else: |
73 | | - if worker_filter.worker_id is not None: |
74 | | - fmt = _STREAM_HANDLER_FMT_WITH_WORKER_ID |
75 | | - else: |
76 | | - fmt = STREAM_HANDLER_FMT |
77 | | - fmt = logging.Formatter(fmt, DATE_FMT) |
| 89 | + match format: |
| 90 | + case LogFormat.JSON: |
| 91 | + fmt = _json_formatter(datefmt=DATE_FMT) |
| 92 | + case LogFormat.LOGFMT: |
| 93 | + fmt = LogFmtFormatter(datefmt=DATE_FMT) |
| 94 | + case LogFormat.DEFAULT: |
| 95 | + if worker_filter.worker_id is not None: |
| 96 | + fmt = _STREAM_HANDLER_FMT_WITH_WORKER_ID |
| 97 | + else: |
| 98 | + fmt = STREAM_HANDLER_FMT |
| 99 | + fmt = logging.Formatter(fmt, DATE_FMT) |
| 100 | + case _: |
| 101 | + raise NotImplementedError(f"invalid log format: {format}") |
78 | 102 | stream_handler.setFormatter(fmt) |
79 | 103 | stream_handler.setLevel(level) |
80 | 104 | stream_handler.addFilter(worker_filter) |
81 | 105 | return [stream_handler] |
82 | 106 |
|
83 | 107 |
|
| 108 | +class LogFmtFormatter(logging.Formatter): |
| 109 | + def format(self, record: logging.LogRecord) -> str: |
| 110 | + logged = dict() |
| 111 | + if record.exc_info and not record.exc_text: |
| 112 | + record.exc_text = self.formatException(record.exc_info) |
| 113 | + logged["exc_info"] = record.exc_text |
| 114 | + for k, v in record.__dict__.items(): |
| 115 | + if k in _LOGGED_ATTRIBUTES and k != "exc_info": |
| 116 | + logged[k] = _encode_value(v) |
| 117 | + return " ".join(f"{k}={v}" for k, v in sorted(logged.items())) |
| 118 | + |
| 119 | + |
| 120 | +def _encode_value(value: Any) -> str: |
| 121 | + if value is None: |
| 122 | + return "" |
| 123 | + if isinstance(value, bool): |
| 124 | + return "true" if value else "false" |
| 125 | + if isinstance(value, numbers.Number): |
| 126 | + return str(value) |
| 127 | + return orjson.dumps(value).decode() |
| 128 | + |
| 129 | + |
84 | 130 | def _json_formatter(datefmt: str) -> BaseJsonFormatter: |
85 | 131 | fmt = OrjsonFormatter( # let's keep logging as fast as possible |
86 | 132 | _LOGGED_ATTRIBUTES, datefmt=datefmt |
|
0 commit comments