1use std::{collections::BTreeMap, fmt::Debug, time::Duration};
4
5use conv::ConvUtil;
6
7use crate::{limiter::Outcome, limits::Sample};
8
9pub trait Aggregator {
14 fn sample(&mut self, sample: Sample) -> Sample;
18
19 #[allow(missing_docs)]
20 fn sample_size(&self) -> usize;
21
22 #[allow(missing_docs)]
23 fn reset(&mut self);
24}
25
26#[derive(Debug)]
28pub struct Average {
29 latency_sum: Duration,
30 in_flight_sum: u128,
31 overload: Outcome,
32 samples: usize,
33}
34
35pub struct Percentile {
37 percentile: f64,
38 overload: Outcome,
39 num_samples: usize,
40 samples: BTreeMap<Duration, Vec<Sample>>,
41}
42
43impl Aggregator for Average {
44 fn sample(&mut self, sample: Sample) -> Sample {
45 self.latency_sum += sample.latency;
46 self.in_flight_sum += sample.in_flight as u128;
47 self.overload = self.overload.overloaded_or(sample.outcome);
48 self.samples += 1;
49 Sample {
50 in_flight: (self.in_flight_sum / self.samples as u128) as usize,
51 latency: self.latency_sum.div_f64(self.samples as f64),
52 outcome: self.overload,
53 }
54 }
55
56 fn sample_size(&self) -> usize {
57 self.samples
58 }
59
60 fn reset(&mut self) {
61 *self = Self::default();
62 }
63}
64
65impl Default for Average {
66 fn default() -> Self {
67 Self {
68 latency_sum: Duration::ZERO,
69 in_flight_sum: 0,
70 overload: Outcome::Success,
71 samples: 0,
72 }
73 }
74}
75
76impl Percentile {
77 #[allow(missing_docs)]
78 pub fn new(percentile: f64) -> Self {
79 assert!(
80 percentile > 0. && percentile < 1.,
81 "percentiles must be between 0 and 1 exclusive"
82 );
83 Self {
84 percentile,
85 ..Default::default()
86 }
87 }
88
89 fn percentile_sample(&self) -> Option<&Sample> {
90 let index = self.percentile_index();
91
92 index.and_then(|index| {
93 self.samples
94 .iter()
95 .flat_map(|(_, sample)| sample)
96 .nth(index)
97 })
98 }
99
100 fn percentile_index(&self) -> Option<usize> {
101 if self.num_samples == 0 {
102 return None;
103 }
104
105 let float_index = self.num_samples as f64 * self.percentile;
106
107 Some(
108 float_index
109 .ceil()
110 .approx_as::<usize>()
111 .expect("percentile should be < 1")
112 - 1,
113 )
114 }
115}
116
117impl Aggregator for Percentile {
118 fn sample(&mut self, sample: Sample) -> Sample {
119 self.overload = self.overload.overloaded_or(sample.outcome);
120 self.samples.entry(sample.latency).or_default().push(sample);
121 self.num_samples += 1;
122
123 let perc_sample = self
124 .percentile_sample()
125 .expect("Sample should exist at expected index");
126
127 Sample {
128 in_flight: perc_sample.in_flight,
134 latency: perc_sample.latency,
135 outcome: self.overload,
136 }
137 }
138
139 fn sample_size(&self) -> usize {
140 self.num_samples
141 }
142
143 fn reset(&mut self) {
144 *self = Self {
145 percentile: self.percentile,
146 ..Default::default()
147 };
148 }
149}
150
151impl Default for Percentile {
152 fn default() -> Self {
153 Self {
154 percentile: 0.5,
155 samples: BTreeMap::new(),
156 num_samples: 0,
157 overload: Outcome::Success,
158 }
159 }
160}
161
162impl Debug for Percentile {
163 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
164 f.debug_struct("Percentile")
165 .field("percentile", &self.percentile)
166 .field("overload", &self.overload)
167 .field("samples", &self.samples)
168 .field("(aggregated sample)", &self.percentile_sample())
169 .finish()
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176
177 #[tokio::test]
178 async fn average() {
179 let mut aggregator = Average::default();
180
181 aggregator.sample(Sample {
182 in_flight: 1,
183 latency: Duration::from_millis(1),
184 outcome: Outcome::Success,
185 });
186
187 aggregator.sample(Sample {
188 in_flight: 5,
189 latency: Duration::from_millis(3),
190 outcome: Outcome::Overload,
191 });
192
193 let sample = aggregator.sample(Sample {
194 in_flight: 3,
195 latency: Duration::from_millis(5),
196 outcome: Outcome::Success,
197 });
198
199 assert_eq!(
200 sample,
201 Sample {
202 in_flight: 3,
203 latency: Duration::from_millis(3),
204 outcome: Outcome::Overload,
205 }
206 );
207 }
208
209 #[tokio::test]
210 async fn average_reset() {
211 let mut aggregator = Average::default();
212
213 aggregator.sample(Sample {
214 in_flight: 1,
215 latency: Duration::from_millis(1),
216 outcome: Outcome::Success,
217 });
218
219 aggregator.reset();
220
221 let sample = aggregator.sample(Sample {
222 in_flight: 3,
223 latency: Duration::from_millis(5),
224 outcome: Outcome::Success,
225 });
226
227 assert_eq!(
228 sample,
229 Sample {
230 in_flight: 3,
231 latency: Duration::from_millis(5),
232 outcome: Outcome::Success,
233 },
234 "should be equal to new sample after reset"
235 )
236 }
237
238 #[tokio::test]
239 async fn percentile_p01() {
240 let mut aggregator = Percentile::new(0.01);
241
242 aggregator.sample(Sample {
243 in_flight: 5,
244 latency: Duration::from_millis(3),
245 outcome: Outcome::Overload,
246 });
247
248 aggregator.sample(Sample {
249 in_flight: 1,
250 latency: Duration::from_millis(1),
251 outcome: Outcome::Success,
252 });
253
254 let sample = aggregator.sample(Sample {
255 in_flight: 3,
256 latency: Duration::from_millis(5),
257 outcome: Outcome::Success,
258 });
259
260 assert_eq!(
261 sample,
262 Sample {
263 in_flight: 1,
264 latency: Duration::from_millis(1),
265 outcome: Outcome::Overload,
266 }
267 );
268 }
269
270 #[tokio::test]
271 async fn percentile_p99() {
272 let mut aggregator = Percentile::new(0.99);
273
274 aggregator.sample(Sample {
275 in_flight: 5,
276 latency: Duration::from_millis(3),
277 outcome: Outcome::Overload,
278 });
279
280 aggregator.sample(Sample {
281 in_flight: 1,
282 latency: Duration::from_millis(1),
283 outcome: Outcome::Success,
284 });
285
286 let sample = aggregator.sample(Sample {
287 in_flight: 3,
288 latency: Duration::from_millis(5),
289 outcome: Outcome::Success,
290 });
291
292 assert_eq!(
293 sample,
294 Sample {
295 in_flight: 3,
296 latency: Duration::from_millis(5),
297 outcome: Outcome::Overload,
298 }
299 );
300 }
301
302 #[tokio::test]
303 async fn percentile_reset() {
304 let mut aggregator = Percentile::new(0.99);
305
306 aggregator.sample(Sample {
307 in_flight: 1,
308 latency: Duration::from_millis(1),
309 outcome: Outcome::Success,
310 });
311
312 aggregator.reset();
313
314 let sample = aggregator.sample(Sample {
315 in_flight: 3,
316 latency: Duration::from_millis(5),
317 outcome: Outcome::Success,
318 });
319
320 assert_eq!(
321 sample,
322 Sample {
323 in_flight: 3,
324 latency: Duration::from_millis(5),
325 outcome: Outcome::Success,
326 },
327 "should be equal to new sample after reset"
328 );
329
330 assert_eq!(
331 aggregator.percentile, 0.99,
332 "percentile shouldn't change after reset"
333 );
334 }
335}