ForcePilot/backend/package/yuxi/channel/runtime/trace_ctx.py

52 lines
1.4 KiB
Python
Raw Normal View History

from __future__ import annotations
import contextvars
import logging
import uuid
from collections.abc import Generator
from contextlib import contextmanager
from typing import Any
_trace_id_var: contextvars.ContextVar[str | None] = contextvars.ContextVar(
"channel_trace_id", default=None
)
class TraceIdFilter(logging.Filter):
"""日志过滤器:自动将当前 trace_id 注入 LogRecord 的 trace_id 属性。
配合 logging.Formatter 使用在日志格式中包含 %(trace_id)s 即可
"""
def filter(self, record: logging.LogRecord) -> bool:
trace_id = get_trace_id()
record.trace_id = trace_id or "-"
return True
def generate_trace_id() -> str:
return uuid.uuid4().hex[:16]
def get_trace_id() -> str | None:
return _trace_id_var.get()
def set_trace_id(trace_id: str | None) -> None:
_trace_id_var.set(trace_id)
@contextmanager
def trace_context(trace_id: str | None = None, **metadata: Any) -> Generator[str, None, None]:
token = _trace_id_var.set(trace_id or generate_trace_id())
try:
yield _trace_id_var.get()
finally:
_trace_id_var.reset(token)
def install_trace_filter(logger_name: str | None = None) -> None:
root_logger = logging.getLogger(logger_name)
if not any(isinstance(f, TraceIdFilter) for f in root_logger.filters):
root_logger.addFilter(TraceIdFilter())