saluki_common/task/
instrument.rs

1use std::{
2    future::Future,
3    pin::Pin,
4    task::{Context, Poll},
5};
6
7use pin_project::pin_project;
8use saluki_metrics::{static_metrics, Counter, Histogram};
9
10#[static_metrics(prefix = runtime_task, labels(task_name))]
11#[derive(Clone)]
12struct Telemetry {
13    #[metric(level = debug)]
14    poll_count: Counter,
15    #[metric(level = trace)]
16    poll_duration_seconds: Histogram,
17}
18
19/// Helper trait for instrumenting futures that are run as asynchronous tasks.
20pub trait TaskInstrument {
21    /// Instruments the future, tracking task-specific metrics about its execution.
22    ///
23    /// Whenever the resulting future is polled, two internal metrics are updated: `runtime_task.poll_count` is
24    /// incremented by one, and `runtime_task.poll_duration_seconds` records the duration of the poll operation, in
25    /// seconds. Both metrics are tagged with the task name provided here (as `task_name:<task name>`).
26    ///
27    /// In general, a unique task name should be provided where possible. If multiple tasks share the same task name,
28    /// they will all update the same metric, which will simply influence the resulting percentiles and make it more
29    /// difficult to isolate outlier poll durations.
30    fn with_task_instrumentation(self, task_name: String) -> InstrumentedTask<Self>
31    where
32        Self: Sized;
33}
34
35impl<F> TaskInstrument for F
36where
37    F: Future + Send + 'static,
38{
39    fn with_task_instrumentation(self, task_name: String) -> InstrumentedTask<Self> {
40        InstrumentedTask {
41            telemetry: Telemetry::new(task_name),
42            inner: self,
43        }
44    }
45}
46
47/// An instrumented task future.
48///
49/// This wraps a `Future` and emits telemetry about the duration of each poll operation.
50#[pin_project]
51pub struct InstrumentedTask<F> {
52    telemetry: Telemetry,
53
54    #[pin]
55    inner: F,
56}
57
58impl<F> Future for InstrumentedTask<F>
59where
60    F: Future,
61{
62    type Output = F::Output;
63
64    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
65        let this = self.project();
66
67        let poll_start = std::time::Instant::now();
68        let result = this.inner.poll(cx);
69        let poll_duration = poll_start.elapsed();
70
71        this.telemetry.poll_count.increment(1);
72        this.telemetry.poll_duration_seconds.record(poll_duration.as_secs_f64());
73
74        result
75    }
76}
77
78#[cfg(test)]
79mod tests {
80    use std::{
81        future::Future,
82        pin::Pin,
83        task::{Context, Poll, Waker},
84    };
85
86    use saluki_metrics::test::TestRecorder;
87
88    use super::*;
89
90    /// A future that returns `Pending` a fixed number of times before completing, letting a test
91    /// drive a known number of polls.
92    struct PendsThenReady {
93        pending_polls: usize,
94    }
95
96    impl Future for PendsThenReady {
97        type Output = ();
98
99        fn poll(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> {
100            if self.pending_polls == 0 {
101                Poll::Ready(())
102            } else {
103                self.pending_polls -= 1;
104                Poll::Pending
105            }
106        }
107    }
108
109    #[test]
110    fn poll_records_one_duration_sample_per_poll() {
111        let recorder = TestRecorder::default();
112        let _guard = metrics::set_default_local_recorder(&recorder);
113
114        // Two `Pending` polls followed by one `Ready` poll: three polls total.
115        let task = PendsThenReady { pending_polls: 2 }.with_task_instrumentation("poll_duration_test".to_string());
116        let mut task = Box::pin(task);
117
118        let waker = Waker::noop();
119        let mut cx = Context::from_waker(waker);
120
121        let mut polls = 0;
122        loop {
123            polls += 1;
124            if task.as_mut().poll(&mut cx).is_ready() {
125                break;
126            }
127        }
128        assert_eq!(polls, 3);
129
130        // The documented behavior is that every poll records one `poll_duration_seconds` sample,
131        // tagged with the task name provided to `with_task_instrumentation`.
132        let samples = recorder
133            .histogram((
134                Telemetry::poll_duration_seconds_name(),
135                &[("task_name", "poll_duration_test")],
136            ))
137            .expect("poll-duration histogram should be registered");
138        assert_eq!(samples.len(), 3, "each poll must record one duration sample");
139        assert!(
140            samples.iter().all(|&sample| sample >= 0.0),
141            "recorded poll durations must be non-negative"
142        );
143
144        // Every poll also increments `poll_count`, tagged with the same task name, so after three
145        // polls the counter reads three.
146        assert_eq!(
147            recorder.counter((Telemetry::poll_count_name(), &[("task_name", "poll_duration_test")])),
148            Some(3)
149        );
150    }
151}