Skip to main content

rig_core/test_utils/
trace_capture.rs

1//! A `tracing` layer that records spans and events for test assertions.
2//!
3//! ```
4//! use rig_core::test_utils::TraceCapture;
5//!
6//! let capture = TraceCapture::default();
7//! tracing::subscriber::with_default(capture.subscriber(), || {
8//!     let _span = tracing::info_span!("work", answer = 42_u64).entered();
9//! });
10//! assert_eq!(capture.spans()[0].u64("answer"), Some(42));
11//! ```
12
13use std::sync::{Arc, Mutex, PoisonError};
14
15use serde_json::{Map, Value};
16use tracing::field::{Field, Visit};
17use tracing::span::{Attributes, Id, Record};
18use tracing::{Event, Level, Subscriber};
19use tracing_subscriber::layer::{Context, Layer, SubscriberExt};
20use tracing_subscriber::registry::LookupSpan;
21
22/// Records every span and event it sees. Clones share one record, so a test
23/// keeps a clone and installs another with [`TraceCapture::subscriber`].
24///
25/// Field values are JSON: strings and `Debug`/`Display` renderings are
26/// strings, integers are numbers and booleans are booleans.
27#[derive(Clone, Default)]
28pub struct TraceCapture(Arc<Mutex<Captured>>);
29
30#[derive(Default)]
31struct Captured {
32    spans: Vec<CapturedSpan>,
33    events: Vec<CapturedEvent>,
34}
35
36impl Captured {
37    fn span_mut(&mut self, id: &Id) -> Option<&mut CapturedSpan> {
38        self.spans
39            .iter_mut()
40            .rev()
41            .find(|span| span.id == id.into_u64())
42    }
43}
44
45/// One span as [`TraceCapture`] saw it.
46#[derive(Clone, Debug)]
47pub struct CapturedSpan {
48    /// The span's id, unique while the span is open.
49    pub id: u64,
50    /// The span's name.
51    pub name: &'static str,
52    /// The span's target.
53    pub target: &'static str,
54    /// The explicit parent, or the current span for a contextual span.
55    pub parent: Option<u64>,
56    /// The parent's name.
57    pub parent_name: Option<&'static str>,
58    /// Every field the span declares, in declaration order.
59    pub declared: Vec<&'static str>,
60    /// The values given when the span was created.
61    pub initial: Map<String, Value>,
62    /// Every value recorded later, in order, repeats included.
63    pub recorded: Vec<(String, Value)>,
64    /// The ids this span follows from.
65    pub follows_from: Vec<u64>,
66}
67
68impl CapturedSpan {
69    /// The latest value of `field`: its last recording, else its initial value.
70    pub fn value(&self, field: &str) -> Option<&Value> {
71        self.recorded
72            .iter()
73            .rev()
74            .find(|(name, _)| name == field)
75            .map(|(_, value)| value)
76            .or_else(|| self.initial.get(field))
77    }
78
79    /// [`Self::value`] as text: a string as it is, anything else as JSON.
80    pub fn text(&self, field: &str) -> Option<String> {
81        self.value(field).map(text)
82    }
83
84    /// The latest value of `field` when it is an unsigned integer.
85    pub fn u64(&self, field: &str) -> Option<u64> {
86        self.value(field).and_then(Value::as_u64)
87    }
88
89    /// Every value of `field` recorded after creation, in order, as text.
90    pub fn recorded_texts(&self, field: &str) -> Vec<String> {
91        self.recorded
92            .iter()
93            .filter(|(name, _)| name == field)
94            .map(|(_, value)| text(value))
95            .collect()
96    }
97
98    /// How many times `field` was recorded after creation.
99    pub fn record_count(&self, field: &str) -> usize {
100        self.recorded
101            .iter()
102            .filter(|(name, _)| name == field)
103            .count()
104    }
105
106    /// The initial values with every later recording applied in order.
107    pub fn values(&self) -> Map<String, Value> {
108        let mut values = self.initial.clone();
109        values.extend(self.recorded.iter().cloned());
110        values
111    }
112
113    /// The span without its id, which differs between runs: name, target,
114    /// parent name, declared fields and [`Self::values`].
115    pub fn summary(&self) -> Value {
116        serde_json::json!({
117            "name": self.name,
118            "target": self.target,
119            "parent": self.parent_name,
120            "fields": self.declared,
121            "values": self.values(),
122        })
123    }
124}
125
126/// One event as [`TraceCapture`] saw it.
127#[derive(Clone, Debug)]
128pub struct CapturedEvent {
129    /// The event's level.
130    pub level: Level,
131    /// The event's target.
132    pub target: &'static str,
133    /// The event's fields, the message under `message`.
134    pub fields: Map<String, Value>,
135}
136
137impl CapturedEvent {
138    /// The event's message.
139    pub fn message(&self) -> String {
140        self.fields.get("message").map(text).unwrap_or_default()
141    }
142}
143
144fn text(value: &Value) -> String {
145    match value {
146        Value::String(text) => text.clone(),
147        other => other.to_string(),
148    }
149}
150
151impl TraceCapture {
152    /// A registry with this capture as its only layer.
153    pub fn subscriber(&self) -> impl Subscriber + Send + Sync + 'static {
154        tracing_subscriber::registry().with(self.clone())
155    }
156
157    /// Every span seen so far, in creation order.
158    pub fn spans(&self) -> Vec<CapturedSpan> {
159        self.lock().spans.clone()
160    }
161
162    /// Every event seen so far, in order.
163    pub fn events(&self) -> Vec<CapturedEvent> {
164        self.lock().events.clone()
165    }
166
167    /// The last span opened so far.
168    pub fn last_span(&self) -> Option<CapturedSpan> {
169        self.lock().spans.last().cloned()
170    }
171
172    /// Every value `field` took on any span, at creation or later, span by
173    /// span in creation order.
174    pub fn values_of(&self, field: &str) -> Vec<Value> {
175        let captured = self.lock();
176        let mut values = Vec::new();
177        for span in &captured.spans {
178            values.extend(span.initial.get(field).cloned());
179            values.extend(
180                span.recorded
181                    .iter()
182                    .filter(|(name, _)| name == field)
183                    .map(|(_, value)| value.clone()),
184            );
185        }
186        values
187    }
188
189    /// The events at WARN, in order, each as its message followed by
190    /// ` name=value` for every other field.
191    pub fn warnings(&self) -> Vec<String> {
192        self.lock()
193            .events
194            .iter()
195            .filter(|event| event.level == Level::WARN)
196            .map(|event| {
197                let mut rendered = event.message();
198                for (name, value) in event.fields.iter().filter(|(name, _)| *name != "message") {
199                    rendered.push_str(&format!(" {name}={}", text(value)));
200                }
201                rendered
202            })
203            .collect()
204    }
205
206    /// Forgets every span and event seen so far.
207    pub fn clear(&self) {
208        let mut captured = self.lock();
209        captured.spans.clear();
210        captured.events.clear();
211    }
212
213    fn lock(&self) -> std::sync::MutexGuard<'_, Captured> {
214        self.0.lock().unwrap_or_else(PoisonError::into_inner)
215    }
216}
217
218impl<S> Layer<S> for TraceCapture
219where
220    S: Subscriber + for<'a> LookupSpan<'a>,
221{
222    fn on_new_span(&self, attrs: &Attributes<'_>, id: &Id, ctx: Context<'_, S>) {
223        let parent = match attrs.parent() {
224            Some(parent) => ctx.span(parent),
225            None if attrs.is_contextual() => ctx.lookup_current(),
226            None => None,
227        };
228        let mut initial = Map::new();
229        attrs.record(&mut Values(&mut initial));
230        let metadata = attrs.metadata();
231        self.lock().spans.push(CapturedSpan {
232            id: id.into_u64(),
233            name: metadata.name(),
234            target: metadata.target(),
235            parent: parent.as_ref().map(|span| span.id().into_u64()),
236            parent_name: parent.as_ref().map(|span| span.name()),
237            declared: metadata.fields().iter().map(|field| field.name()).collect(),
238            initial,
239            recorded: Vec::new(),
240            follows_from: Vec::new(),
241        });
242    }
243
244    fn on_record(&self, id: &Id, values: &Record<'_>, _: Context<'_, S>) {
245        let mut fields = Map::new();
246        values.record(&mut Values(&mut fields));
247        if let Some(span) = self.lock().span_mut(id) {
248            span.recorded.extend(fields);
249        }
250    }
251
252    fn on_follows_from(&self, id: &Id, follows: &Id, _: Context<'_, S>) {
253        if let Some(span) = self.lock().span_mut(id) {
254            span.follows_from.push(follows.into_u64());
255        }
256    }
257
258    fn on_event(&self, event: &Event<'_>, _: Context<'_, S>) {
259        let mut fields = Map::new();
260        event.record(&mut Values(&mut fields));
261        self.lock().events.push(CapturedEvent {
262            level: *event.metadata().level(),
263            target: event.metadata().target(),
264            fields,
265        });
266    }
267}
268
269struct Values<'a>(&'a mut Map<String, Value>);
270
271impl Visit for Values<'_> {
272    fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
273        self.0
274            .insert(field.name().into(), Value::String(format!("{value:?}")));
275    }
276
277    fn record_str(&mut self, field: &Field, value: &str) {
278        self.0.insert(field.name().into(), Value::from(value));
279    }
280
281    fn record_u64(&mut self, field: &Field, value: u64) {
282        self.0.insert(field.name().into(), Value::from(value));
283    }
284
285    fn record_i64(&mut self, field: &Field, value: i64) {
286        self.0.insert(field.name().into(), Value::from(value));
287    }
288
289    fn record_bool(&mut self, field: &Field, value: bool) {
290        self.0.insert(field.name().into(), Value::from(value));
291    }
292}