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)