Skip to main content

gwr_track/
perfetto_trace_builder.rs

1// Copyright (c) 2025 Graphcore Ltd. All rights reserved.
2
3//! Support for generating Perfetto traces.
4//!
5//! The public API allows the creation of various TrackDescriptors and
6//! TrackEvents required to represent events from the [crate::tracker::Track]
7//! trait, each contained within their own timestamped TracePacket.
8//!
9//! To create trace files that can be opened using the Perfetto UI the
10//! TracePackets must be wrapped in a Trace message. The build_trace_to_bytes()
11//! function is provided to do this and serialise the data ready for writing.
12//!
13//! Multiple TracePackets can included within a single Trace message and
14//! multiple trace messages can be written consecutively to the same Perfetto
15//! trace file.
16
17use std::collections::HashMap;
18
19use gwr_perfetto::protos::trace_packet::Data;
20use gwr_perfetto::protos::track_descriptor::StaticOrDynamicName;
21use gwr_perfetto::protos::{
22    CounterDescriptor, Trace, TracePacket, TrackDescriptor, TrackEvent, counter_descriptor,
23    trace_packet, track_event,
24};
25use prost::Message;
26use rand::random;
27
28use crate::Id;
29
30/// State for a trace builder instance.
31pub struct PerfettoTraceBuilder {
32    trusted_packet_sequence_id: u32,
33    id_to_name: HashMap<u64, String>,
34}
35
36impl Default for PerfettoTraceBuilder {
37    fn default() -> Self {
38        Self {
39            trusted_packet_sequence_id: random(),
40            id_to_name: HashMap::new(),
41        }
42    }
43}
44
45impl PerfettoTraceBuilder {
46    /// Create a Perfetto trace builder.
47    ///
48    /// Each trace builder instance will use a unique TrustedPacketSequenceId
49    /// when creating packets.
50    #[must_use]
51    pub fn new() -> Self {
52        PerfettoTraceBuilder::default()
53    }
54
55    fn build_incremental_counter_track_descriptor(
56        &mut self,
57        id: Id,
58        parent: Id,
59        name: &str,
60    ) -> TrackDescriptor {
61        let counter_desc = CounterDescriptor {
62            r#type: Some(counter_descriptor::BuiltinCounterType::CounterUnspecified as i32),
63            unit: Some(counter_descriptor::Unit::Count as i32),
64            is_incremental: Some(true),
65            ..Default::default()
66        };
67
68        let mut track_descriptor = self.build_track_descriptor(id, parent, name);
69        track_descriptor.counter = Some(counter_desc);
70
71        track_descriptor
72    }
73
74    fn build_absolute_counter_track_descriptor(
75        &mut self,
76        id: Id,
77        parent: Id,
78        name: &str,
79    ) -> TrackDescriptor {
80        let counter_desc = CounterDescriptor {
81            r#type: Some(counter_descriptor::BuiltinCounterType::CounterUnspecified as i32),
82            unit: Some(counter_descriptor::Unit::Count as i32),
83            is_incremental: Some(false),
84            ..Default::default()
85        };
86
87        let mut track_descriptor = self.build_track_descriptor(id, parent, name);
88        track_descriptor.counter = Some(counter_desc);
89
90        track_descriptor
91    }
92
93    fn build_track_descriptor(&mut self, id: Id, parent: Id, name: &str) -> TrackDescriptor {
94        self.set_id_to_name(id, name);
95
96        TrackDescriptor {
97            uuid: Some(id.0),
98            parent_uuid: Some(parent.0),
99            static_or_dynamic_name: Some(StaticOrDynamicName::AtraceName(name.to_string())),
100            ..Default::default()
101        }
102    }
103
104    fn set_id_to_name(&mut self, id: Id, name: &str) {
105        self.id_to_name.insert(id.0, name.to_owned());
106    }
107
108    fn build_incremental_counter_track_event(
109        &self,
110        id: Id,
111        other: Id,
112        increment: i64,
113    ) -> TrackEvent {
114        let mut track_event = self.build_track_event(id, other);
115        track_event.set_type(track_event::Type::Counter);
116        track_event.counter_value_field =
117            Some(track_event::CounterValueField::CounterValue(increment));
118
119        track_event
120    }
121
122    fn build_absolute_double_counter_track_event(&self, id: Id, absolute: f64) -> TrackEvent {
123        let mut track_event = self.build_track_event(id, Id(0));
124        track_event.set_type(track_event::Type::Counter);
125        track_event.counter_value_field =
126            Some(track_event::CounterValueField::DoubleCounterValue(absolute));
127
128        track_event
129    }
130
131    fn build_track_event(&self, id: Id, other: Id) -> TrackEvent {
132        TrackEvent {
133            track_uuid: Some(id.0),
134            name_field: Some(track_event::NameField::Name(self.id_to_name(id, other))),
135            ..Default::default()
136        }
137    }
138
139    fn id_to_name(&self, id: Id, other: Id) -> String {
140        let name = match id.0 {
141            0 => "root",
142            _ => match self.id_to_name.get(&other.0) {
143                Some(name) => name,
144                None => "UNKNOWN",
145            },
146        };
147
148        name.to_string()
149    }
150
151    /// Build a TracePacket containing the TrackDescriptor for
152    /// [crate::tracker::Track::enter] and [crate::tracker::Track::exit] events
153    /// (using an incremental counter).
154    #[must_use]
155    pub fn build_enter_exit_track_descriptor_trace_packet(
156        &mut self,
157        current_time_ns: u64,
158        id: Id,
159        parent: Id,
160        name: &str,
161    ) -> TracePacket {
162        let track_descriptor = self.build_incremental_counter_track_descriptor(id, parent, name);
163
164        self.build_track_descriptor_trace_packet(current_time_ns, track_descriptor)
165    }
166
167    /// Build a TracePacket containing the TrackDescriptor for a sequence of
168    /// [crate::tracker::Track::value]s.
169    #[must_use]
170    pub fn build_value_track_descriptor_trace_packet(
171        &mut self,
172        current_time_ns: u64,
173        id: Id,
174        parent: Id,
175        name: &str,
176    ) -> TracePacket {
177        let track_descriptor = self.build_absolute_counter_track_descriptor(id, parent, name);
178
179        self.build_track_descriptor_trace_packet(current_time_ns, track_descriptor)
180    }
181
182    /// Build a TracePacket containing the TrackDescriptor for activity slices
183    /// denoted by [crate::tracker::Track::begin_activity] and
184    /// [crate::tracker::Track::end_activity] events.
185    #[must_use]
186    pub fn build_activity_track_descriptor_trace_packet(
187        &mut self,
188        current_time_ns: u64,
189        id: Id,
190        parent: Id,
191        name: &str,
192    ) -> TracePacket {
193        let track_descriptor = self.build_track_descriptor(id, parent, name);
194
195        self.build_track_descriptor_trace_packet(current_time_ns, track_descriptor)
196    }
197
198    fn build_track_descriptor_trace_packet(
199        &self,
200        current_time_ns: u64,
201        track_descriptor: TrackDescriptor,
202    ) -> TracePacket {
203        let mut trace_packet = self.build_trace_packet(current_time_ns);
204        trace_packet.data = Some(Data::TrackDescriptor(track_descriptor));
205
206        trace_packet
207    }
208
209    /// Build a TracePacket containing a TrackEvent to represent an
210    /// [crate::tracker::Track::enter] event (as an incremental counter
211    /// update).
212    #[must_use]
213    pub fn build_enter_track_event_trace_packet(
214        &self,
215        current_time_ns: u64,
216        id: Id,
217        other: Id,
218    ) -> TracePacket {
219        let track_event = self.build_incremental_counter_track_event(id, other, 1);
220
221        self.build_track_event_trace_packet(current_time_ns, track_event)
222    }
223
224    /// Build a TracePacket containing a TrackEvent to represent an
225    /// [crate::tracker::Track::exit] event (as an incremental counter
226    /// update).
227    #[must_use]
228    pub fn build_exit_track_event_trace_packet(
229        &self,
230        current_time_ns: u64,
231        id: Id,
232        other: Id,
233    ) -> TracePacket {
234        let track_event = self.build_incremental_counter_track_event(id, other, -1);
235
236        self.build_track_event_trace_packet(current_time_ns, track_event)
237    }
238
239    /// Build a TracePacket containing the TrackEvent for a floating point
240    /// [crate::tracker::Track::value].
241    #[must_use]
242    pub fn build_value_track_event_trace_packet(
243        &self,
244        current_time_ns: u64,
245        id: Id,
246        value: f64,
247    ) -> TracePacket {
248        let track_event = self.build_absolute_double_counter_track_event(id, value);
249
250        self.build_track_event_trace_packet(current_time_ns, track_event)
251    }
252
253    /// Build a TracePacket containing the TrackEvent for a
254    /// [crate::tracker::Track::begin_activity] SliceBegin event.
255    #[must_use]
256    pub fn build_activity_begin_trace_packet(
257        &self,
258        current_time_ns: u64,
259        id: Id,
260        name: &str,
261        correlation_id: Option<u64>,
262    ) -> TracePacket {
263        let track_event = build_slice_track_event(
264            id,
265            Some(name),
266            track_event::Type::SliceBegin,
267            correlation_id,
268        );
269
270        self.build_track_event_trace_packet(current_time_ns, track_event)
271    }
272
273    /// Build a TracePacket containing the TrackEvent for a
274    /// [crate::tracker::Track::end_activity] SliceEnd event.
275    #[must_use]
276    pub fn build_activity_end_trace_packet(&self, current_time_ns: u64, id: Id) -> TracePacket {
277        let track_event = build_slice_track_event(id, None, track_event::Type::SliceEnd, None);
278
279        self.build_track_event_trace_packet(current_time_ns, track_event)
280    }
281
282    fn build_track_event_trace_packet(
283        &self,
284        current_time_ns: u64,
285        track_event: TrackEvent,
286    ) -> TracePacket {
287        let mut trace_packet = self.build_trace_packet(current_time_ns);
288        trace_packet.data = Some(trace_packet::Data::TrackEvent(track_event));
289
290        trace_packet
291    }
292
293    fn build_trace_packet(&self, current_time_ns: u64) -> TracePacket {
294        TracePacket {
295            timestamp: Some(current_time_ns),
296            optional_trusted_packet_sequence_id: Some(
297                trace_packet::OptionalTrustedPacketSequenceId::TrustedPacketSequenceId(
298                    self.trusted_packet_sequence_id,
299                ),
300            ),
301            ..Default::default()
302        }
303    }
304
305    /// Build a Trace message containing the passed TracePackets and serialise
306    /// it to unsigned bytes.
307    #[must_use]
308    pub fn build_trace_to_bytes(&self, trace_packets: Vec<TracePacket>) -> Vec<u8> {
309        PerfettoTraceBuilder::build_trace(trace_packets).encode_to_vec()
310    }
311
312    fn build_trace(trace_packets: Vec<TracePacket>) -> Trace {
313        Trace {
314            packet: trace_packets,
315        }
316    }
317}
318
319fn build_slice_track_event(
320    id: Id,
321    name: Option<&str>,
322    event_type: track_event::Type,
323    correlation_id: Option<u64>,
324) -> TrackEvent {
325    let mut track_event = match name {
326        Some(name) => build_named_track_event(id, name),
327        None => TrackEvent {
328            track_uuid: Some(id.0),
329            ..Default::default()
330        },
331    };
332    track_event.set_type(event_type);
333    if let Some(correlation_id) = correlation_id {
334        track_event.correlation_id_field = Some(track_event::CorrelationIdField::CorrelationId(
335            correlation_id,
336        ));
337    }
338    track_event
339}
340
341fn build_named_track_event(id: Id, name: &str) -> TrackEvent {
342    TrackEvent {
343        track_uuid: Some(id.0),
344        name_field: Some(track_event::NameField::Name(name.to_string())),
345        ..Default::default()
346    }
347}
348
349#[cfg(test)]
350mod tests {
351    use gwr_perfetto::protos::trace_packet::Data;
352
353    use super::*;
354
355    #[test]
356    fn activity_packets_are_perfetto_slices() {
357        let mut builder = PerfettoTraceBuilder::new();
358        let descriptor =
359            builder.build_activity_track_descriptor_trace_packet(0, Id(11), Id(10), "pe::op");
360        let begin = builder.build_activity_begin_trace_packet(42, Id(11), "add_task (add)", None);
361        let correlated_begin =
362            builder.build_activity_begin_trace_packet(43, Id(11), "add compute", Some(99));
363        let end = builder.build_activity_end_trace_packet(84, Id(11));
364
365        let Some(Data::TrackDescriptor(descriptor)) = descriptor.data else {
366            panic!("expected activity track descriptor");
367        };
368        assert_eq!(descriptor.uuid, Some(11));
369        assert_eq!(descriptor.parent_uuid, Some(10));
370        assert!(descriptor.counter.is_none());
371
372        let Some(Data::TrackEvent(begin)) = begin.data else {
373            panic!("expected activity begin track event");
374        };
375        assert_eq!(begin.track_uuid, Some(11));
376        assert_eq!(begin.r#type, Some(track_event::Type::SliceBegin as i32));
377        assert_eq!(
378            begin.name_field,
379            Some(track_event::NameField::Name("add_task (add)".to_string()))
380        );
381
382        let Some(Data::TrackEvent(correlated_begin)) = correlated_begin.data else {
383            panic!("expected correlated activity begin track event");
384        };
385        assert_eq!(correlated_begin.track_uuid, Some(11));
386        assert_eq!(
387            correlated_begin.correlation_id_field,
388            Some(track_event::CorrelationIdField::CorrelationId(99))
389        );
390
391        let Some(Data::TrackEvent(end)) = end.data else {
392            panic!("expected activity end track event");
393        };
394        assert_eq!(end.track_uuid, Some(11));
395        assert_eq!(end.r#type, Some(track_event::Type::SliceEnd as i32));
396    }
397}