v0.166.0
  1"""
  2Capture log records emitted during a block.
  3
  4The framework's own loggers don't propagate to the root logger, so a
  5root-attached handler sees nothing. `capture_logs` attaches directly to the
  6loggers you name (by default the whole `plain` and `app` trees) and takes them
  7down again on exit.
  8"""
  9
 10import logging
 11from collections.abc import Generator
 12from contextlib import contextmanager
 13from typing import TYPE_CHECKING
 14
 15from .captured import Captured
 16
 17if TYPE_CHECKING:
 18    from opentelemetry.trace import SpanContext
 19
 20__all__ = ["CapturedLogs", "capture_logs"]
 21
 22# The trees Plain code logs into. `plain.mcp`, `plain.jobs` and friends
 23# propagate up to `plain`, so naming the root of each tree catches them all.
 24DEFAULT_LOGGER_NAMES = ("plain", "app")
 25
 26
 27class _RecordingHandler(logging.Handler):
 28    """Records each emit alongside the span context current at that moment.
 29
 30    The span context is kept beside the record rather than set on it. A
 31    LogRecord's `__dict__` is how Plain carries structured context (see
 32    `plain.logs.formatters`), so an extra attribute here would show up as a
 33    key=value pair in every other handler's output — including the OTel
 34    LoggingHandler's exported attributes.
 35    """
 36
 37    def __init__(self) -> None:
 38        super().__init__(level=logging.DEBUG)
 39        self.entries: list[tuple[logging.LogRecord, SpanContext]] = []
 40        # One handler can be attached to several loggers at once (the default
 41        # is two). A record that reaches more than one of them is still one
 42        # record, so it's counted once.
 43        self._seen: set[int] = set()
 44
 45    def emit(self, record: logging.LogRecord) -> None:
 46        if id(record) in self._seen:
 47            return
 48        self._seen.add(id(record))
 49
 50        from opentelemetry.trace import get_current_span
 51
 52        self.entries.append((record, get_current_span().get_span_context()))
 53
 54
 55class CapturedLogs(Captured[logging.LogRecord]):
 56    """
 57    The records logged during a `capture_logs` block, in the order they were
 58    logged. Each one is a `logging.LogRecord`.
 59    """
 60
 61    def __init__(self) -> None:
 62        super().__init__(helper="capture_logs")
 63        # By the record's id. The record itself can't carry it: see
 64        # _RecordingHandler.
 65        self._span_context_of: dict[int, SpanContext] = {}
 66
 67    def finish_with_span_contexts(
 68        self, entries: list[tuple[logging.LogRecord, SpanContext]]
 69    ) -> None:
 70        """Say what was captured: each record, with the span context that
 71        was current when it was logged."""
 72        self._span_context_of = {id(record): context for record, context in entries}
 73        self.finish(record for record, _ in entries)
 74
 75    @property
 76    def messages(self) -> list[str]:
 77        """The formatted message of every captured record."""
 78        return [record.getMessage() for record in self]
 79
 80    def __repr__(self) -> str:
 81        if not self.finished:
 82            return super().__repr__()
 83        return f"<CapturedLogs {self.messages!r}>"
 84
 85    def span_context_for(self, message: str) -> SpanContext:
 86        """
 87        The OpenTelemetry span context that was current when the record with
 88        this message was emitted.
 89
 90            with capture_spans() as spans, capture_logs() as logs:
 91                do_the_thing()
 92
 93            [span] = spans.filter(name="claim job")
 94            assert logs.span_context_for("Claim failed").trace_id == (
 95                span.context.trace_id
 96            )
 97
 98        A log emitted with no span current has an invalid (all-zero) context,
 99        which is what an exporter would ship — that's the failure this is
100        usually asserting against. Raises if the message wasn't logged exactly
101        once, so a typo or a duplicate can't pass silently.
102        """
103        matches = [
104            self._span_context_of[id(record)]
105            for record in self
106            if record.getMessage() == message
107        ]
108        if not matches:
109            raise LookupError(
110                f"No log record with message {message!r}. Captured: {self.messages!r}"
111            )
112        if len(matches) > 1:
113            raise LookupError(
114                f"{len(matches)} log records with message {message!r} — "
115                "span_context_for needs exactly one."
116            )
117        return matches[0]
118
119
120@contextmanager
121def capture_logs(
122    *logger_names: str, level: int = logging.DEBUG
123) -> Generator[CapturedLogs]:
124    """
125    The log records emitted during the block.
126
127        with capture_logs() as logs:
128            Client().get("/boom/")
129
130        assert "Server error" in logs.messages
131        assert logs[0].path == "/boom/"
132
133    With no arguments this captures the `plain` and `app` trees. Name loggers
134    to narrow it:
135
136        with capture_logs("plain.jobs") as logs:
137            ...
138
139    Each logger's level is lowered for the block and restored on exit, as is
140    any global `logging.disable()` in effect — otherwise a level set elsewhere
141    could swallow the records under test.
142    """
143    names = logger_names or DEFAULT_LOGGER_NAMES
144    handler = _RecordingHandler()
145    loggers = [logging.getLogger(name) for name in names]
146
147    original_levels = [logger.level for logger in loggers]
148    original_disable = logging.root.manager.disable
149
150    for logger in loggers:
151        logger.addHandler(handler)
152        logger.setLevel(level)
153    logging.root.manager.disable = 0
154    captured = CapturedLogs()
155    try:
156        yield captured
157    finally:
158        logging.root.manager.disable = original_disable
159        for logger, original_level in zip(loggers, original_levels, strict=True):
160            logger.removeHandler(handler)
161            logger.setLevel(original_level)
162        captured.finish_with_span_contexts(handler.entries)