1"""
2Capture the OpenTelemetry spans and metrics emitted during a block.
3
4The global tracer/meter providers are install-once per process, so the two
5install helpers are idempotent — repeated calls return the same source, and
6every capture reads from that one.
7
8The OpenTelemetry SDK imports are deferred into the install helpers so that
9importing `plain.testing` (e.g. for `Client`) doesn't pay the SDK import cost.
10"""
11
12from collections.abc import Generator, Mapping, Sequence
13from contextlib import contextmanager
14from typing import TYPE_CHECKING, cast
15
16from .captured import Captured, CaptureSource
17
18if TYPE_CHECKING:
19 from opentelemetry.sdk.metrics.export import (
20 DataPointT,
21 HistogramDataPoint,
22 Metric,
23 NumberDataPoint,
24 )
25 from opentelemetry.sdk.trace import ReadableSpan, TracerProvider
26 from opentelemetry.sdk.trace.export import SpanExportResult
27 from opentelemetry.trace import SpanKind
28 from opentelemetry.util.types import AttributeValue
29
30__all__ = ["CapturedMetrics", "CapturedSpans", "capture_metrics", "capture_spans"]
31
32# Where every capture of each kind reads from. Made the first time one is
33# asked for, along with the provider that feeds it.
34_span_source: CaptureSource[ReadableSpan] | None = None
35_metric_source: CaptureSource[Metric] | None = None
36
37# The tracer provider installed here, for capturing. The process has one
38# global provider and it can be set once, so whoever needs one for capturing
39# gets it from `tracer_provider_for_capturing()`.
40_tracer_provider_for_capturing: TracerProvider | None = None
41
42
43def tracer_provider_for_capturing() -> TracerProvider:
44 """
45 The process's tracer provider, as long as it is one installed here.
46
47 Installs it the first time it's asked for, when nothing has installed
48 one, and returns the same provider after that. It starts with no span
49 processors: it records nothing until a capture adds its own.
50
51 Raises when the provider in place was installed by something else.
52 `set_tracer_provider` is one-shot, so a second one would be ignored
53 without a word, every capture would come up empty, and what the tests
54 do would be exported to wherever that provider sends it.
55 """
56 global _tracer_provider_for_capturing
57
58 from opentelemetry import trace
59 from opentelemetry.sdk.trace import TracerProvider
60
61 if _tracer_provider_for_capturing is None:
62 if not isinstance(trace.get_tracer_provider(), trace.ProxyTracerProvider):
63 raise RuntimeError(
64 "A global tracer provider is already installed, and not by"
65 " plain.testing, so spans can't be captured: whatever installed"
66 " it has to leave it out of a test run. (plain.connect does."
67 " It reads PLAIN_TEST_RUNNING, which `plain test` sets.)"
68 )
69 provider = TracerProvider()
70 trace.set_tracer_provider(provider)
71 _tracer_provider_for_capturing = provider
72
73 return _tracer_provider_for_capturing
74
75
76def _install_test_tracer() -> CaptureSource[ReadableSpan]:
77 global _span_source
78 if _span_source is None:
79 from opentelemetry.sdk.trace.export import (
80 SimpleSpanProcessor,
81 SpanExportResult,
82 )
83 from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
84 InMemorySpanExporter,
85 )
86
87 class SpanExporterForCaptures(InMemorySpanExporter):
88 """
89 Keeps the spans that end while a capture is open. Once installed,
90 the provider hands this every span for the rest of the run, and
91 the ones nobody is capturing would otherwise be kept until the
92 next capture emptied them.
93 """
94
95 def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
96 if not source.capturing:
97 return SpanExportResult.SUCCESS
98 return super().export(spans)
99
100 exporter = SpanExporterForCaptures()
101 source = CaptureSource(read=exporter.get_finished_spans, clear=exporter.clear)
102
103 # Raises when the provider in place isn't ours to add to.
104 provider = tracer_provider_for_capturing()
105 provider.add_span_processor(SimpleSpanProcessor(exporter))
106 _span_source = source
107 return _span_source
108
109
110def _install_test_meter() -> CaptureSource[Metric]:
111 global _metric_source
112 if _metric_source is None:
113 from opentelemetry import metrics
114 from opentelemetry.sdk.metrics import (
115 Counter,
116 Histogram,
117 MeterProvider,
118 UpDownCounter,
119 )
120 from opentelemetry.sdk.metrics.export import (
121 AggregationTemporality,
122 InMemoryMetricReader,
123 )
124
125 # Delta temporality so each collection only reports what happened
126 # since the last one — that's what makes the drain-on-entry in
127 # capture_metrics() actually isolate one test's metrics from the
128 # counters accumulated by everything that ran before it.
129 reader = InMemoryMetricReader(
130 preferred_temporality={
131 Counter: AggregationTemporality.DELTA,
132 UpDownCounter: AggregationTemporality.DELTA,
133 Histogram: AggregationTemporality.DELTA,
134 }
135 )
136 provider = MeterProvider(metric_readers=[reader])
137 metrics.set_meter_provider(provider)
138 if metrics.get_meter_provider() is not provider:
139 raise RuntimeError(
140 "A global meter provider is already installed, and not by"
141 " plain.testing, so metrics can't be captured: whatever"
142 " installed it has to leave it out of a test run."
143 " (plain.connect does. It reads PLAIN_TEST_RUNNING, which"
144 " `plain test` sets.)"
145 )
146
147 # The reader hands over what was recorded since it was last asked,
148 # and then no longer has it. So what it hands over is kept here, for
149 # every open capture to read.
150 collected: list[Metric] = []
151
152 def collect() -> list[Metric]:
153 """Collect what the reader holds — which also asks every
154 observable instrument for its current value — and keep it."""
155 data = reader.get_metrics_data()
156 if data is not None:
157 for resource_metrics in data.resource_metrics:
158 for scope_metrics in resource_metrics.scope_metrics:
159 collected.extend(scope_metrics.metrics)
160 return collected
161
162 _metric_source = CaptureSource(read=collect, clear=collected.clear)
163 return _metric_source
164
165
166class CapturedSpans(Captured["ReadableSpan"]):
167 """
168 The spans that ended during a `capture_spans` block, in the order they
169 ended. Each one is OpenTelemetry's own `ReadableSpan`.
170 """
171
172 def __init__(self) -> None:
173 super().__init__(helper="capture_spans")
174
175 def filter(
176 self, *, name: str | None = None, kind: SpanKind | None = None
177 ) -> list[ReadableSpan]:
178 """
179 The spans with this name, of this kind, or both.
180
181 [span] = spans.filter(name="claim job")
182 server_spans = spans.filter(kind=SpanKind.SERVER)
183 """
184 if name is None and kind is None:
185 raise TypeError("filter() needs a name=, a kind=, or both")
186 return [
187 span
188 for span in self
189 if (name is None or span.name == name)
190 and (kind is None or span.kind == kind)
191 ]
192
193
194class CapturedMetrics(Captured["Metric"]):
195 """
196 The metrics collected for a `capture_metrics` block, in the order they
197 were collected. Each one is OpenTelemetry's own `Metric`.
198
199 What a test usually wants are a metric's data points, and those come in
200 two kinds: `number_points(name)` for a counter or a gauge, and
201 `histogram_points(name)` for a histogram.
202 """
203
204 def __init__(self) -> None:
205 super().__init__(helper="capture_metrics")
206
207 def number_points(
208 self, name: str, *, attributes: Mapping[str, AttributeValue] | None = None
209 ) -> list[NumberDataPoint]:
210 """
211 The data points of the counter, up-down counter or gauge with this
212 name. Each has a `value`.
213
214 points = metrics.number_points(
215 "messaging.client.consumed.messages",
216 attributes={"plain.jobs.outcome": "lost"},
217 )
218 assert sum(point.value for point in points) == 1
219
220 Pass `attributes` to keep only the points that carry all of them.
221 """
222 from opentelemetry.sdk.metrics.export import NumberDataPoint
223
224 points = self._points(name, attributes)
225 for point in points:
226 if not isinstance(point, NumberDataPoint):
227 raise TypeError(
228 f"{name!r} is a histogram — read it with"
229 f" `histogram_points({name!r})`."
230 )
231 return cast("list[NumberDataPoint]", points)
232
233 def histogram_points(
234 self, name: str, *, attributes: Mapping[str, AttributeValue] | None = None
235 ) -> list[HistogramDataPoint]:
236 """
237 The data points of the histogram with this name. Each has a `count`,
238 a `sum`, a `min` and a `max`.
239
240 points = metrics.histogram_points(
241 "db.client.response.returned_rows",
242 attributes={"db.operation.name": "SELECT"},
243 )
244 assert sum(point.sum for point in points) == 5
245
246 Pass `attributes` to keep only the points that carry all of them.
247 """
248 from opentelemetry.sdk.metrics.export import HistogramDataPoint
249
250 points = self._points(name, attributes)
251 for point in points:
252 if not isinstance(point, HistogramDataPoint):
253 raise TypeError(
254 f"{name!r} is not a histogram — read it with"
255 f" `number_points({name!r})`."
256 )
257 return cast("list[HistogramDataPoint]", points)
258
259 def _points(
260 self, name: str, attributes: Mapping[str, AttributeValue] | None
261 ) -> list[DataPointT]:
262 wanted = attributes or {}
263 points = []
264 for metric in self:
265 if metric.name != name:
266 continue
267 for point in metric.data.data_points:
268 carried = point.attributes or {}
269 if all(
270 key in carried and carried[key] == value
271 for key, value in wanted.items()
272 ):
273 points.append(point)
274 return points
275
276
277@contextmanager
278def capture_spans() -> Generator[CapturedSpans]:
279 """
280 The OpenTelemetry spans that end during the block.
281
282 with capture_spans() as spans:
283 Client().get("/")
284
285 [server_span] = spans.filter(kind=SpanKind.SERVER)
286 """
287 source = _install_test_tracer()
288 captured = CapturedSpans()
289 with source.capturing_into(captured):
290 yield captured
291
292
293@contextmanager
294def capture_metrics() -> Generator[CapturedMetrics]:
295 """
296 The OpenTelemetry metrics recorded during the block.
297
298 with capture_metrics() as metrics:
299 Client().get("/")
300
301 assert metrics.histogram_points("http.server.request.duration")
302
303 Metrics are collected when the block ends, which is also when an
304 observable instrument (a gauge that reports a pool's size, say) is asked
305 for its value. So whatever it observes has to still be there: open the
306 capture inside the block that keeps it alive, not around it.
307 """
308 source = _install_test_meter()
309 captured = CapturedMetrics()
310 # What was recorded before the block is collected as the capture starts,
311 # so it lands before the point this capture reads from.
312 with source.capturing_into(captured):
313 yield captured