saluki_common/task/
instrument.rs1use 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
19pub trait TaskInstrument {
21 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#[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 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 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 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 assert_eq!(
147 recorder.counter((Telemetry::poll_count_name(), &[("task_name", "poll_duration_test")])),
148 Some(3)
149 );
150 }
151}