quiche/recovery/gcongestion/bbr2/rtt_jump_detector/
hmm.rs1use std::time::Duration;
28use std::time::Instant;
29
30use super::RttJumpUpdate;
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46enum HmmState {
47 Normal,
48 Transient,
49 Persistent,
50}
51
52impl HmmState {
53 fn as_index(self) -> usize {
54 match self {
55 HmmState::Normal => 0,
56 HmmState::Transient => 1,
57 HmmState::Persistent => 2,
58 }
59 }
60
61 fn from_index(index: usize) -> Self {
62 match index {
63 0 => HmmState::Normal,
64 1 => HmmState::Transient,
65 2 => HmmState::Persistent,
66 _ => unreachable!("invalid HMM state index"),
67 }
68 }
69
70 fn from_episode(episode: HmmEpisode) -> Self {
71 match episode {
72 HmmEpisode::Idle => HmmState::Normal,
73 HmmEpisode::Active => HmmState::Transient,
74 HmmEpisode::Persistent => HmmState::Persistent,
75 }
76 }
77
78 fn to_episode(self) -> HmmEpisode {
79 match self {
80 HmmState::Normal => HmmEpisode::Idle,
81 HmmState::Transient => HmmEpisode::Active,
82 HmmState::Persistent => HmmEpisode::Persistent,
83 }
84 }
85
86 fn is_more_elevated_than(self, other: Self) -> bool {
87 self.as_index() > other.as_index()
88 }
89}
90
91#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
92enum HmmEpisode {
93 #[default]
94 Idle,
95 Active,
96 Persistent,
97}
98
99const HMM_STATE_COUNT: usize = 3;
100const HMM_BINS: usize = 5;
101
102const HMM_WARMUP_SAMPLES: u32 = 8;
104
105const HMM_MIN_DWELL: Duration = Duration::from_millis(80);
107
108const HMM_TRANSIENT_DWELL_RTTS: u32 = 1;
111
112const HMM_PERSIST_DWELL_RTTS: u32 = 5;
115
116const HMM_CLEAR_DWELL_RTTS: u32 = 1;
119
120const HMM_PERSIST_CONFIRM_SAMPLES: u32 = 3;
122
123const HMM_TRANSITION: [[f32; HMM_STATE_COUNT]; HMM_STATE_COUNT] = [
132 [0.9939, 0.0055, 0.0006],
135 [0.4500, 0.3500, 0.2000],
137 [0.0243, 0.1603, 0.8154],
139];
140
141const HMM_EMISSION: [[f32; HMM_BINS]; HMM_STATE_COUNT] = [
144 [0.9996, 0.0001, 0.0001, 0.0001, 0.0001],
146 [0.1000, 0.3500, 0.3000, 0.1500, 0.1000],
148 [0.0001, 0.5000, 0.2200, 0.1600, 0.1199],
150];
151
152const HMM_STD_BIN_EDGES: [f32; HMM_BINS - 1] = [1.0, 2.0, 4.0, 8.0];
155
156const HMM_DISP_EWMA_ALPHA: f32 = 0.0625;
158
159const HMM_ELEVATION_FLOOR: f32 = 0.20;
162
163fn hmm_std_bin(elevation: f32, edges: &[f32; HMM_BINS - 1]) -> usize {
165 let mut bin = 0;
166 while bin < edges.len() && elevation >= edges[bin] {
167 bin += 1;
168 }
169 bin
170}
171
172fn hmm_forward_step(
174 alpha: &mut [f32; HMM_STATE_COUNT], bin: usize,
175 transition: &[[f32; HMM_STATE_COUNT]; HMM_STATE_COUNT],
176 emission: &[[f32; HMM_BINS]; HMM_STATE_COUNT],
177) {
178 let mut next = [0.0f32; HMM_STATE_COUNT];
179 for (j, next_j) in next.iter_mut().enumerate() {
180 let mut predicted = 0.0f32;
181 for (i, &alpha_i) in alpha.iter().enumerate() {
182 predicted += alpha_i * transition[i][j];
183 }
184 *next_j = predicted * emission[j][bin];
185 }
186
187 let sum: f32 = next.iter().sum();
188 if sum > 0.0 {
189 let scale = 1.0 / sum;
190 for (dst, src) in alpha.iter_mut().zip(next.iter()) {
191 *dst = src * scale;
192 }
193 } else {
194 *alpha = [1.0, 0.0, 0.0];
195 }
196}
197
198fn hmm_argmax(alpha: &[f32; HMM_STATE_COUNT]) -> HmmState {
200 let mut best = 0;
201 for state in 1..HMM_STATE_COUNT {
202 if alpha[state] > alpha[best] {
203 best = state;
204 }
205 }
206 HmmState::from_index(best)
207}
208
209fn relative_elevation(rtt_sample: Duration, baseline: Option<Duration>) -> f32 {
210 match baseline {
211 Some(baseline) if !baseline.is_zero() =>
212 (rtt_sample.as_secs_f32() / baseline.as_secs_f32() - 1.0).max(0.0),
213 _ => 0.0,
214 }
215}
216
217fn dwell_for_transition(
218 current: HmmState, pending: HmmState, rtt_sample: Duration,
219) -> Duration {
220 let rtts = if pending.is_more_elevated_than(current) {
221 match pending {
222 HmmState::Persistent => HMM_PERSIST_DWELL_RTTS,
223 HmmState::Transient => HMM_TRANSIENT_DWELL_RTTS,
224 HmmState::Normal => unreachable!("normal is not elevated"),
225 }
226 } else {
227 HMM_CLEAR_DWELL_RTTS
228 };
229
230 HMM_MIN_DWELL.max(rtt_sample * rtts)
231}
232
233#[derive(Debug)]
234pub(super) struct HmmDetector {
235 episode: HmmEpisode,
236 alpha: [f32; HMM_STATE_COUNT],
238 valid_samples: u32,
240 warmup_contrib: u32,
242 dispersion: f32,
245 baseline: Option<Duration>,
247 pending_state: HmmState,
249 pending_start_time: Option<Instant>,
251 pending_samples: u32,
253 episode_counted: bool,
255}
256
257impl Default for HmmDetector {
258 fn default() -> Self {
259 Self {
260 episode: HmmEpisode::Idle,
261 alpha: [1.0, 0.0, 0.0],
262 valid_samples: 0,
263 warmup_contrib: 0,
264 dispersion: 0.0,
265 baseline: None,
266 pending_state: HmmState::Normal,
267 pending_start_time: None,
268 pending_samples: 0,
269 episode_counted: false,
270 }
271 }
272}
273
274impl HmmDetector {
275 pub(super) fn on_rtt_sample(
278 &mut self, rtt_sample: Duration, event_time: Instant,
279 full_bandwidth_reached: bool,
280 ) -> RttJumpUpdate {
281 self.valid_samples = self.valid_samples.saturating_add(1);
282
283 let rel = relative_elevation(rtt_sample, self.baseline);
284
285 if self.valid_samples <= HMM_WARMUP_SAMPLES {
286 if rel > 0.0 {
287 self.warmup_contrib = self.warmup_contrib.saturating_add(1);
288 self.dispersion +=
289 (rel - self.dispersion) / self.warmup_contrib as f32;
290 }
291 self.update_baseline(rtt_sample);
292 return RttJumpUpdate::None;
293 }
294
295 let scale = self.dispersion.max(HMM_ELEVATION_FLOOR);
296 let elevation = (rel - self.dispersion).max(0.0) / scale;
297 let bin = hmm_std_bin(elevation, &HMM_STD_BIN_EDGES);
298 hmm_forward_step(&mut self.alpha, bin, &HMM_TRANSITION, &HMM_EMISSION);
299
300 let raw = hmm_argmax(&self.alpha);
301
302 if !full_bandwidth_reached {
303 self.reset_pending(HmmState::Normal);
304 self.update_adaptive_state(rtt_sample, false);
305 return RttJumpUpdate::None;
306 }
307
308 let committed_state = HmmState::from_episode(self.episode);
309 let mut committed_to_normal = false;
310 let mut update = RttJumpUpdate::None;
311
312 if raw == committed_state {
313 self.reset_pending(committed_state);
314 } else {
315 if raw != self.pending_state || self.pending_start_time.is_none() {
316 self.pending_state = raw;
317 self.pending_start_time = Some(event_time);
318 self.pending_samples = 1;
319 } else {
320 self.pending_samples = self.pending_samples.saturating_add(1);
321 }
322
323 let episode_start = self.pending_start_time.unwrap_or(event_time);
324 let dwell = dwell_for_transition(committed_state, raw, rtt_sample);
325 let samples_ok = raw != HmmState::Persistent ||
326 self.pending_samples >= HMM_PERSIST_CONFIRM_SAMPLES;
327 let time_ok =
328 event_time.saturating_duration_since(episode_start) >= dwell;
329
330 if samples_ok && time_ok {
331 match raw {
332 HmmState::Persistent => {
333 if !self.episode_counted {
334 self.episode_counted = true;
335 update = RttJumpUpdate::PersistentConfirmed {
336 episode_start_time: episode_start,
337 };
338 }
339 self.episode = raw.to_episode();
340 },
341 HmmState::Transient => {
342 self.episode = raw.to_episode();
343 },
344 HmmState::Normal => {
345 committed_to_normal = true;
346 self.episode_counted = false;
347 self.episode = raw.to_episode();
348 },
349 }
350 self.reset_pending(raw);
351 }
352 }
353
354 self.update_adaptive_state(rtt_sample, committed_to_normal);
355 update
356 }
357
358 fn reset_pending(&mut self, state: HmmState) {
359 self.pending_state = state;
360 self.pending_start_time = None;
361 self.pending_samples = 0;
362 }
363
364 fn update_adaptive_state(
365 &mut self, rtt_sample: Duration, committed_to_normal: bool,
366 ) {
367 if committed_to_normal {
368 self.baseline = Some(rtt_sample);
369 } else {
370 self.update_baseline(rtt_sample);
371 }
372
373 let rel = relative_elevation(rtt_sample, self.baseline);
374
375 if matches!(self.episode, HmmEpisode::Idle) &&
376 self.pending_start_time.is_none()
377 {
378 self.dispersion += HMM_DISP_EWMA_ALPHA * (rel - self.dispersion);
379 }
380 }
381
382 fn update_baseline(&mut self, rtt_sample: Duration) {
383 self.baseline =
384 Some(self.baseline.map_or(rtt_sample, |min| min.min(rtt_sample)));
385 }
386
387 #[cfg(test)]
388 pub(super) fn is_rtt_jump_active(&self) -> bool {
389 !matches!(self.episode, HmmEpisode::Idle)
390 }
391
392 #[cfg(test)]
393 pub(super) fn is_rtt_jump_persistent(&self) -> bool {
394 matches!(self.episode, HmmEpisode::Persistent)
395 }
396}
397
398#[cfg(test)]
399#[path = "hmm_tests.rs"]
400mod tests;