1use std::cmp::{max, min};
4use std::collections::VecDeque;
5use std::sync::atomic::{AtomicBool, Ordering};
6use web_time_compat::{Duration, Instant};
7
8use super::params::RoundTripEstimatorParams;
9use super::{CongestionWindow, State};
10
11use thiserror::Error;
12use tor_error::{ErrorKind, HasKind};
13
14#[derive(Error, Debug, Clone)]
16#[non_exhaustive]
17pub(crate) enum Error {
18 #[error("Informed of a SENDME we weren't expecting")]
21 MismatchedEstimationCall,
22}
23
24impl HasKind for Error {
25 fn kind(&self) -> ErrorKind {
26 use Error as E;
27 match self {
28 E::MismatchedEstimationCall => ErrorKind::TorProtocolViolation,
29 }
30 }
31}
32
33#[derive(Debug)]
35#[allow(dead_code)]
36pub(crate) struct RoundtripTimeEstimator {
37 sendme_expected_from: VecDeque<Instant>,
45 last_rtt: Option<Duration>,
49 ewma_rtt: Option<Duration>,
53 min_rtt: Option<Duration>,
57 max_rtt: Option<Duration>,
61 params: RoundTripEstimatorParams,
63 clock_stalled: AtomicBool,
66}
67
68#[allow(dead_code)]
69impl RoundtripTimeEstimator {
70 pub(crate) fn new(params: &RoundTripEstimatorParams) -> Self {
73 Self {
74 sendme_expected_from: Default::default(),
75 last_rtt: None,
76 ewma_rtt: None,
77 min_rtt: None,
78 max_rtt: None,
79 params: params.clone(),
80 clock_stalled: AtomicBool::default(),
81 }
82 }
83
84 pub(crate) fn is_ready(&self) -> bool {
86 !self.clock_stalled() && self.last_rtt.is_some()
87 }
88
89 pub(crate) fn clock_stalled(&self) -> bool {
91 self.clock_stalled.load(Ordering::SeqCst)
92 }
93
94 pub(crate) fn ewma_rtt(&self) -> Option<Duration> {
96 self.ewma_rtt
97 }
98
99 pub(crate) fn min_rtt(&self) -> Option<Duration> {
101 self.min_rtt
102 }
103
104 pub(crate) fn max_rtt(&self) -> Option<Duration> {
106 self.max_rtt
107 }
108
109 pub(crate) fn expect_sendme(&mut self, now: Instant) {
112 self.sendme_expected_from.push_back(now);
113 }
114
115 fn can_crosscheck_with_current_estimate(&self, in_slow_start: bool) -> bool {
121 !in_slow_start && self.ewma_rtt.is_some()
125 }
126
127 fn is_clock_stalled(&self, raw_rtt: Duration, in_slow_start: bool) -> bool {
130 if raw_rtt.is_zero() {
131 self.clock_stalled.store(true, Ordering::SeqCst);
133 true
134 } else if self.can_crosscheck_with_current_estimate(in_slow_start) {
135 let ewma_rtt = self
136 .ewma_rtt
137 .expect("ewma_rtt was not checked by can_crosscheck_with_current_estimate?!");
138
139 const DELTA_DISCREPANCY_RATIO_MAX: u32 = 5000;
143 if raw_rtt > ewma_rtt * DELTA_DISCREPANCY_RATIO_MAX {
145 true
152 } else if ewma_rtt > raw_rtt * DELTA_DISCREPANCY_RATIO_MAX {
153 self.clock_stalled.load(Ordering::SeqCst)
156 } else {
157 self.clock_stalled.store(false, Ordering::SeqCst);
159 false
160 }
161 } else {
162 false
164 }
165 }
166
167 pub(crate) fn update(
181 &mut self,
182 now: Instant,
183 state: &State,
184 cwnd: &CongestionWindow,
185 ) -> Result<ClockStall, Error> {
186 let data_sent_at = self
187 .sendme_expected_from
188 .pop_front()
189 .ok_or(Error::MismatchedEstimationCall)?;
190 let raw_rtt = now.saturating_duration_since(data_sent_at);
191
192 if self.is_clock_stalled(raw_rtt, state.in_slow_start()) {
193 return Ok(ClockStall::Detected);
194 }
195
196 self.max_rtt = self.max_rtt.max(Some(raw_rtt));
197 self.last_rtt = Some(raw_rtt);
198
199 let ewma_n = u64::from(if state.in_slow_start() {
201 self.params.ewma_ss_max()
202 } else {
203 min(
204 (cwnd.update_rate(state) * (self.params.ewma_cwnd_pct().as_percent())) / 100,
205 self.params.ewma_max(),
206 )
207 });
208 let ewma_n = max(ewma_n, 2);
209
210 let raw_rtt_usec = raw_rtt.as_micros() as u64;
212 let prev_ewma_rtt_usec = self.ewma_rtt.map(|rtt| rtt.as_micros() as u64);
213
214 let new_ewma_rtt_usec = match prev_ewma_rtt_usec {
222 None => raw_rtt_usec,
223 Some(prev_ewma_rtt_usec) => {
224 ((raw_rtt_usec * 2) + ((ewma_n - 1) * prev_ewma_rtt_usec)) / (ewma_n + 1)
225 }
226 };
227 let ewma_rtt = Duration::from_micros(new_ewma_rtt_usec);
228 self.ewma_rtt = Some(ewma_rtt);
229
230 let Some(min_rtt) = self.min_rtt else {
231 self.min_rtt = self.ewma_rtt;
232 return Ok(ClockStall::NotDetected);
233 };
234
235 if cwnd.get() == cwnd.min() && !state.in_slow_start() {
236 let max = max(ewma_rtt, min_rtt).as_micros() as u64;
238 let min = min(ewma_rtt, min_rtt).as_micros() as u64;
239 let rtt_reset_pct = u64::from(self.params.rtt_reset_pct().as_percent());
240 let min_rtt = Duration::from_micros(
241 (rtt_reset_pct * max / 100) + (100 - rtt_reset_pct) * min / 100,
242 );
243
244 self.min_rtt = Some(min_rtt);
245 } else if self.ewma_rtt < self.min_rtt {
246 self.min_rtt = self.ewma_rtt;
247 }
248
249 Ok(ClockStall::NotDetected)
250 }
251}
252
253#[derive(Copy, Clone, Debug, PartialEq, Eq)]
255pub(crate) enum ClockStall {
256 Detected,
258 NotDetected,
260}
261
262#[cfg(test)]
263mod test {
264 #![allow(clippy::bool_assert_comparison)]
266 #![allow(clippy::clone_on_copy)]
267 #![allow(clippy::dbg_macro)]
268 #![allow(clippy::mixed_attributes_style)]
269 #![allow(clippy::print_stderr)]
270 #![allow(clippy::print_stdout)]
271 #![allow(clippy::single_char_pattern)]
272 #![allow(clippy::unwrap_used)]
273 #![allow(clippy::unchecked_time_subtraction)]
274 #![allow(clippy::useless_vec)]
275 #![allow(clippy::needless_pass_by_value)]
276 #![allow(clippy::string_slice)] use web_time_compat::{Duration, Instant, InstantExt};
280
281 use crate::congestion::test_utils::{new_cwnd, new_rtt_estimator};
282
283 use super::*;
284
285 #[derive(Debug)]
286 struct RttTestSample {
287 sent_usec_in: u64,
288 sendme_received_usec_in: u64,
289 cwnd_in: u32,
290 ss_in: bool,
291 last_rtt_usec_out: u64,
292 ewma_rtt_usec_out: u64,
293 min_rtt_usec_out: u64,
294 }
295
296 impl From<[u64; 7]> for RttTestSample {
297 fn from(arr: [u64; 7]) -> Self {
298 Self {
299 sent_usec_in: arr[0],
300 sendme_received_usec_in: arr[1],
301 cwnd_in: arr[2] as u32,
302 ss_in: arr[3] == 1,
303 last_rtt_usec_out: arr[4],
304 ewma_rtt_usec_out: arr[5],
305 min_rtt_usec_out: arr[6],
306 }
307 }
308 }
309 impl RttTestSample {
310 fn test(&self, estimator: &mut RoundtripTimeEstimator, start: Instant) {
311 let state = if self.ss_in {
312 State::SlowStart
313 } else {
314 State::Steady
315 };
316 let mut cwnd = new_cwnd();
317 cwnd.set(self.cwnd_in);
318 let sent = start + Duration::from_micros(self.sent_usec_in);
319 let sendme_received = start + Duration::from_micros(self.sendme_received_usec_in);
320
321 estimator.expect_sendme(sent);
322 estimator
323 .update(sendme_received, &state, &cwnd)
324 .expect("Error on RTT update");
325 assert_eq!(
326 estimator.last_rtt,
327 Some(Duration::from_micros(self.last_rtt_usec_out))
328 );
329 assert_eq!(
330 estimator.ewma_rtt,
331 Some(Duration::from_micros(self.ewma_rtt_usec_out))
332 );
333 assert_eq!(
334 estimator.min_rtt,
335 Some(Duration::from_micros(self.min_rtt_usec_out))
336 );
337 }
338 }
339
340 #[test]
341 fn test_vectors() {
342 let mut rtt = new_rtt_estimator();
343 let now = Instant::get();
344 let vectors = [
346 [100000, 200000, 124, 1, 100000, 100000, 100000],
347 [200000, 300000, 124, 1, 100000, 100000, 100000],
348 [350000, 500000, 124, 1, 150000, 133333, 100000],
349 [500000, 550000, 124, 1, 50000, 77777, 77777],
350 [600000, 700000, 124, 1, 100000, 92592, 77777],
351 [700000, 750000, 124, 1, 50000, 64197, 64197],
352 [750000, 875000, 124, 0, 125000, 104732, 104732],
353 [875000, 900000, 124, 0, 25000, 51577, 104732],
354 [900000, 950000, 200, 0, 50000, 50525, 50525],
355 ];
356 for vect in vectors {
357 let vect = RttTestSample::from(vect);
358 eprintln!("Testing vector: {:?}", vect);
359 vect.test(&mut rtt, now);
360 }
361 }
362}