1use std::{cell::Cell, fmt};
2
3#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash)]
11enum Phase {
12 Active,
14 FiltersStopping(Manner),
16 TransportShutdown(Manner),
18 Stopped(Manner),
24}
25
26impl Phase {
27 fn stage(self) -> u8 {
29 match self {
30 Phase::Active => 0,
31 Phase::FiltersStopping(_) => 1,
32 Phase::TransportShutdown(_) => 2,
33 Phase::Stopped(_) => 3,
34 }
35 }
36
37 fn reached(self, other: Phase) -> bool {
39 self.stage() >= other.stage()
40 }
41
42 fn manner(self) -> Manner {
44 match self {
45 Phase::FiltersStopping(m) | Phase::TransportShutdown(m) | Phase::Stopped(m) => m,
46 Phase::Active => Manner::Graceful,
47 }
48 }
49}
50
51#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
55enum Manner {
56 Graceful,
58 Terminating,
60 ForceClosed,
62}
63
64pub struct Flags {
65 bits: Cell<FlagsKind>,
66 phase: Cell<Phase>,
67}
68
69bitflags::bitflags! {
70 #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
71 pub struct FlagsKind: u16 {
72 const RD_PAUSED = 1 << 1;
74 const RD_BACKPRESSURE = 1 << 2;
76 const RD_EOF = 1 << 3;
78
79 const RD_NOTIFY = 1 << 4;
81 const RD_NOTIFIED = 1 << 5;
83
84 const RD_WR_BACKPRESSURE = 1 << 6;
86 const RD_FILTER_PAUSED = 1 << 7;
88
89 const BUF_R_READY = 1 << 8;
91
92 const WR_FLUSH = 1 << 9;
94 const WR_PAUSED = 1 << 10;
96 const WR_SEND_OP = 1 << 11;
98 const WR_FILTER_PAUSED = 1 << 12;
100
101 const DSP_TIMEOUT = 1 << 13;
103 const DSP_W_BACKPRESSURE = 1 << 14;
105 const DIRECT_WR_SUP = 1 << 15;
107 }
108}
109
110impl Clone for Flags {
111 fn clone(&self) -> Self {
112 Self {
113 bits: Cell::new(self.bits.get()),
114 phase: Cell::new(self.phase.get()),
115 }
116 }
117}
118
119impl Flags {
120 pub(crate) fn new(direct_wr: bool) -> Self {
121 let bits = if direct_wr {
122 FlagsKind::WR_PAUSED | FlagsKind::DIRECT_WR_SUP
123 } else {
124 FlagsKind::WR_PAUSED
125 };
126 Self {
127 bits: Cell::new(bits),
128 phase: Cell::new(Phase::Active),
129 }
130 }
131
132 pub(crate) fn new_stopped() -> Self {
133 Self {
134 bits: Cell::new(FlagsKind::empty()),
135 phase: Cell::new(Phase::Stopped(Manner::Graceful)),
136 }
137 }
138
139 fn contains(&self, f: FlagsKind) -> bool {
140 self.bits.get().contains(f)
141 }
142
143 fn intersects(&self, f: FlagsKind) -> bool {
144 self.bits.get().intersects(f)
145 }
146
147 fn insert(&self, f: FlagsKind) {
148 let mut flags = self.bits.get();
149 flags.insert(f);
150 self.bits.set(flags);
151 }
152
153 fn remove(&self, f: FlagsKind) {
154 let mut flags = self.bits.get();
155 flags.remove(f);
156 self.bits.set(flags);
157 }
158
159 fn reached(&self, phase: Phase) -> bool {
160 self.phase.get().reached(phase)
161 }
162
163 fn ending_at_least(&self, manner: Manner) -> bool {
164 self.phase.get().manner() >= manner
165 }
166
167 pub(crate) fn is_peer_gone(&self) -> bool {
173 self.is_closed()
174 || self.reached(Phase::TransportShutdown(Manner::Graceful))
175 || self.is_terminating()
176 }
177
178 pub(crate) fn is_closed(&self) -> bool {
180 matches!(self.phase.get(), Phase::Stopped(_))
181 }
182
183 pub(crate) fn is_aborted(&self) -> bool {
189 self.is_closed() || self.is_terminating()
190 }
191
192 pub(crate) fn is_stopping(&self) -> bool {
198 self.reached(Phase::TransportShutdown(Manner::Graceful))
199 }
200
201 pub(crate) fn is_terminating(&self) -> bool {
206 self.ending_at_least(Manner::Terminating)
207 }
208
209 pub(crate) fn is_force_closing(&self) -> bool {
218 self.ending_at_least(Manner::ForceClosed)
219 }
220
221 pub(crate) fn is_stopping_or_terminating(&self) -> bool {
222 self.is_stopping() || self.is_terminating()
223 }
224
225 pub(crate) fn is_active(&self) -> bool {
230 self.phase.get() == Phase::Active
231 }
232
233 pub(crate) fn is_stopping_filters(&self) -> bool {
234 self.reached(Phase::FiltersStopping(Manner::Graceful))
235 }
236
237 pub(crate) fn is_write_flush(&self) -> bool {
238 self.intersects(FlagsKind::WR_FLUSH)
239 }
240
241 pub(crate) fn is_shutting_down_filters(&self) -> bool {
247 self.phase.get() == Phase::FiltersStopping(Manner::Graceful)
248 }
249
250 pub(crate) fn is_direct_wr_enabled(&self) -> bool {
251 self.contains(FlagsKind::DIRECT_WR_SUP)
252 }
253
254 pub(crate) fn set_direct_wr_enabled(&self, enabled: bool) {
255 if enabled {
256 self.insert(FlagsKind::DIRECT_WR_SUP);
257 } else {
258 self.remove(FlagsKind::DIRECT_WR_SUP);
259 }
260 }
261
262 pub(crate) fn is_read_paused(&self) -> bool {
263 self.contains(FlagsKind::RD_PAUSED)
264 }
265
266 pub(crate) fn is_write_paused(&self) -> bool {
267 self.contains(FlagsKind::WR_PAUSED)
268 }
269
270 pub(crate) fn is_read_ready(&self) -> bool {
271 self.contains(FlagsKind::BUF_R_READY)
272 }
273
274 pub(crate) fn is_read_notify(&self) -> bool {
275 self.contains(FlagsKind::RD_NOTIFY)
276 }
277
278 #[cfg(test)]
279 pub(crate) fn is_read_notified(&self) -> bool {
280 self.contains(FlagsKind::RD_NOTIFIED)
281 }
282
283 pub(crate) fn is_rd_backpressure(&self) -> bool {
284 self.contains(FlagsKind::RD_BACKPRESSURE)
285 }
286
287 pub(crate) fn is_read_eof(&self) -> bool {
288 self.contains(FlagsKind::RD_EOF)
289 }
290
291 pub(crate) fn is_wr_backpressure(&self) -> bool {
292 self.contains(FlagsKind::DSP_W_BACKPRESSURE)
293 }
294
295 pub(crate) fn is_wr_send_scheduled(&self) -> bool {
296 self.contains(FlagsKind::WR_SEND_OP)
297 }
298
299 pub(crate) fn is_read_wr_backpressure(&self) -> bool {
300 self.contains(FlagsKind::RD_WR_BACKPRESSURE)
301 }
302
303 pub(crate) fn set_read_wr_backpressure(&self) {
304 self.insert(FlagsKind::RD_WR_BACKPRESSURE);
305 }
306
307 pub(crate) fn unset_read_wr_backpressure(&self) {
308 self.remove(FlagsKind::RD_WR_BACKPRESSURE);
309 }
310
311 pub(crate) fn is_read_filter_paused(&self) -> bool {
312 self.contains(FlagsKind::RD_FILTER_PAUSED)
313 }
314
315 pub(crate) fn set_read_filter_paused(&self) {
316 self.insert(FlagsKind::RD_FILTER_PAUSED);
317 }
318
319 pub(crate) fn unset_read_filter_paused(&self) {
320 self.remove(FlagsKind::RD_FILTER_PAUSED);
321 }
322
323 pub(crate) fn is_write_filter_paused(&self) -> bool {
324 self.contains(FlagsKind::WR_FILTER_PAUSED)
325 }
326
327 pub(crate) fn set_write_filter_paused(&self) {
328 self.insert(FlagsKind::WR_FILTER_PAUSED);
329 }
330
331 pub(crate) fn unset_write_filter_paused(&self) {
332 self.remove(FlagsKind::WR_FILTER_PAUSED);
333 }
334
335 pub(crate) fn is_read_paused_or_backpressure(&self) -> bool {
336 self.intersects(FlagsKind::RD_PAUSED | FlagsKind::RD_BACKPRESSURE)
337 }
338
339 pub(crate) fn set_read_paused(&self) {
340 self.insert(FlagsKind::RD_PAUSED);
341 }
342
343 pub(crate) fn set_write_paused(&self) {
344 self.insert(FlagsKind::WR_PAUSED);
345 }
346
347 pub(crate) fn set_wr_send_scheduled(&self) {
348 self.insert(FlagsKind::WR_SEND_OP);
349 }
350
351 pub(crate) fn set_wr_backpressure(&self) {
352 self.insert(FlagsKind::DSP_W_BACKPRESSURE);
353 }
354
355 pub(crate) fn enter_filters_stopping(&self) {
357 if self.phase.get() == Phase::Active {
358 self.phase.set(Phase::FiltersStopping(Manner::Graceful));
359 }
360 }
361
362 pub(crate) fn begin_terminate(&self, force: bool) -> bool {
374 let manner = if force {
375 Manner::ForceClosed
376 } else {
377 Manner::Terminating
378 };
379 let phase = self.phase.get();
380 let started = phase.manner() < Manner::Terminating;
381 self.phase.set(match phase {
382 Phase::Stopped(_) => return false,
383 Phase::Active => Phase::FiltersStopping(manner),
384 Phase::FiltersStopping(m) => Phase::FiltersStopping(m.max(manner)),
385 Phase::TransportShutdown(m) => Phase::TransportShutdown(m.max(manner)),
386 });
387
388 if started {
389 self.insert(FlagsKind::BUF_R_READY);
390 }
391 started
392 }
393
394 pub(crate) fn set_stopped(&self) {
395 let manner = self.phase.get().manner();
396 self.phase.set(Phase::Stopped(manner));
397 }
398
399 pub(crate) fn set_wants_write_flush(&self) {
400 self.insert(FlagsKind::WR_FLUSH);
401 }
402
403 pub(crate) fn set_read_notify(&self) {
404 self.insert(FlagsKind::RD_NOTIFY);
405 }
406
407 pub(crate) fn set_read_notified(&self) {
408 self.insert(FlagsKind::RD_NOTIFIED);
409 }
410
411 pub(crate) fn set_read_ready(&self) {
412 self.insert(FlagsKind::BUF_R_READY);
413 }
414
415 pub(crate) fn set_read_ready_and_backpressure(&self) {
416 self.insert(FlagsKind::RD_PAUSED | FlagsKind::BUF_R_READY | FlagsKind::RD_BACKPRESSURE);
417 }
418
419 pub(crate) fn set_read_eof(&self) {
420 self.insert(FlagsKind::RD_EOF);
421 }
422
423 pub(crate) fn enter_transport_shutdown(&self) {
425 match self.phase.get() {
426 Phase::Active => self.phase.set(Phase::TransportShutdown(Manner::Graceful)),
427 Phase::FiltersStopping(m) => self.phase.set(Phase::TransportShutdown(m)),
428 Phase::TransportShutdown(_) | Phase::Stopped(_) => {}
429 }
430 }
431
432 pub(crate) fn unset_write_paused(&self) {
433 self.remove(FlagsKind::WR_PAUSED);
434 }
435
436 pub(crate) fn unset_wr_send_scheduled(&self) {
437 self.remove(FlagsKind::WR_SEND_OP);
438 }
439
440 pub(crate) fn unset_wr_backpressure(&self) {
441 self.remove(FlagsKind::DSP_W_BACKPRESSURE);
442 }
443
444 pub(crate) fn unset_wr_backpressure_and_flush(&self) {
445 self.remove(FlagsKind::DSP_W_BACKPRESSURE | FlagsKind::WR_FLUSH);
446 }
447
448 pub(crate) fn unset_read_ready_and_backpressure(&self) {
452 self.remove(FlagsKind::BUF_R_READY | FlagsKind::RD_BACKPRESSURE);
453 }
454
455 pub(crate) fn unset_read_ready(&self) {
456 self.remove(FlagsKind::BUF_R_READY);
457 }
458
459 pub(crate) fn unset_read_paused(&self) {
460 self.remove(FlagsKind::RD_PAUSED);
461 }
462
463 pub(crate) fn take_read_notified(&self) -> bool {
467 if self.contains(FlagsKind::RD_NOTIFIED) {
468 self.remove(FlagsKind::RD_NOTIFY | FlagsKind::RD_NOTIFIED);
469 true
470 } else {
471 false
472 }
473 }
474
475 pub(crate) fn check_dispatcher_timeout(&self) -> bool {
477 if self.contains(FlagsKind::DSP_TIMEOUT) {
478 self.remove(FlagsKind::DSP_TIMEOUT);
479 true
480 } else {
481 false
482 }
483 }
484
485 pub(crate) fn check_dispatcher_timeout_unset(&self) -> bool {
487 if self.contains(FlagsKind::DSP_TIMEOUT) {
488 false
489 } else {
490 self.insert(FlagsKind::DSP_TIMEOUT);
491 true
492 }
493 }
494}
495
496impl fmt::Debug for Flags {
497 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
498 write!(f, "{:?} | {:?}", self.bits.get(), self.phase.get())
499 }
500}
501
502#[cfg(test)]
503mod tests {
504 use super::*;
505
506 #[test]
507 fn flags() {
508 assert!(format!("{:?}", Flags::new_stopped()).contains("Stopped"));
509 assert!(format!("{:?}", Flags::new(false)).contains("Active"));
510 assert!(format!("{:?}", FlagsKind::RD_EOF).contains("RD_EOF"));
511 assert_eq!(FlagsKind::RD_EOF, FlagsKind::RD_EOF);
512 assert_ne!(FlagsKind::RD_EOF, FlagsKind::RD_PAUSED);
513 }
514
515 #[test]
516 fn flag_bits_are_distinct() {
517 let all = FlagsKind::all().bits();
518 assert_eq!(
519 all.count_ones() as usize,
520 FlagsKind::all().iter().count(),
521 "two flags share a bit"
522 );
523 }
524
525 #[test]
526 fn shutdown_phase_only_moves_forward() {
527 let f = Flags::new(false);
528 assert!(!f.is_stopping_filters());
529 assert!(!f.is_stopping());
530
531 f.enter_filters_stopping();
532 assert!(f.is_stopping_filters());
533 assert!(f.is_shutting_down_filters());
534 assert!(!f.is_stopping());
535
536 f.enter_transport_shutdown();
537 assert!(f.is_stopping());
538 assert!(f.is_stopping_filters());
540 assert!(!f.is_shutting_down_filters());
541
542 f.enter_filters_stopping();
544 assert!(f.is_stopping());
545
546 f.set_stopped();
548 f.enter_filters_stopping();
549 f.enter_transport_shutdown();
550 assert!(f.is_closed());
551 }
552
553 #[test]
561 fn stopped_is_the_last_state_whichever_route_reached_it() {
562 let graceful = Flags::new(false);
564 graceful.enter_filters_stopping();
565 graceful.enter_transport_shutdown();
566 graceful.set_stopped();
567
568 let hup = Flags::new(false);
570 hup.set_stopped();
571
572 for f in [&graceful, &hup] {
573 assert!(f.is_closed());
574 assert!(!f.is_active());
575 assert!(f.is_stopping());
576 assert!(f.is_stopping_filters());
577 assert!(!f.is_shutting_down_filters());
579 }
580 }
581
582 #[test]
583 fn a_stopped_connection_is_gone_without_having_been_aborted() {
584 let f = Flags::new(false);
585 f.set_stopped();
586 assert!(f.is_closed());
587 assert!(f.is_peer_gone());
588 assert!(f.is_aborted());
589 assert!(!f.is_active());
590 assert!(!f.is_terminating());
592 assert!(!f.is_force_closing());
593 }
594
595 #[test]
596 fn terminate_runs_teardown_once_and_force_still_escalates() {
597 let f = Flags::new(false);
598 assert!(f.begin_terminate(false));
599 assert!(f.is_terminating());
600 assert!(!f.is_force_closing());
601 assert!(f.is_stopping_filters());
603 assert!(f.is_read_ready());
604
605 assert!(!f.begin_terminate(false));
607
608 assert!(!f.begin_terminate(true));
610 assert!(f.is_force_closing());
611 assert!(f.is_terminating());
612 }
613
614 #[test]
615 fn force_close_is_never_downgraded() {
616 let f = Flags::new(false);
617 assert!(f.begin_terminate(true));
618 assert!(f.is_force_closing());
619 f.begin_terminate(false);
620 assert!(
621 f.is_force_closing(),
622 "graceful terminate downgraded an abort"
623 );
624 }
625
626 #[test]
633 fn every_representable_state_is_reachable() {
634 const OPS: usize = 5;
635 let mut seen = std::collections::HashSet::new();
636
637 for len in 1..=5u32 {
638 for mut code in 0..OPS.pow(len) {
639 let f = Flags::new(false);
640 for _ in 0..len {
641 let op = code % OPS;
642 code /= OPS;
643 match op {
644 0 => f.enter_filters_stopping(),
645 1 => {
647 if f.is_stopping_filters() {
648 f.enter_transport_shutdown();
649 }
650 }
651 2 => {
652 f.begin_terminate(false);
653 }
654 3 => {
655 f.begin_terminate(true);
656 }
657 _ => f.set_stopped(),
658 }
659 seen.insert(f.phase.get());
660 }
661 }
662 }
663
664 let manners = [Manner::Graceful, Manner::Terminating, Manner::ForceClosed];
665 let expected: std::collections::HashSet<_> = [Phase::Active]
666 .into_iter()
667 .chain(manners.map(Phase::FiltersStopping))
668 .chain(manners.map(Phase::TransportShutdown))
669 .chain(manners.map(Phase::Stopped))
670 .collect();
671 assert_eq!(seen.len(), 10);
672 assert_eq!(seen, expected);
673 }
674
675 #[test]
679 fn finishing_the_filters_keeps_how_the_connection_ends() {
680 for (force, manner) in [(false, Manner::Terminating), (true, Manner::ForceClosed)] {
681 let f = Flags::new(false);
682 f.begin_terminate(force);
683 assert_eq!(f.phase.get(), Phase::FiltersStopping(manner));
684 f.enter_transport_shutdown();
685 assert_eq!(f.phase.get(), Phase::TransportShutdown(manner));
686 assert_eq!(f.is_force_closing(), force);
687 }
688 }
689
690 #[test]
693 fn a_stopped_connection_remembers_how_it_ended() {
694 for (force, manner) in [(false, Manner::Terminating), (true, Manner::ForceClosed)] {
695 let f = Flags::new(false);
697 f.begin_terminate(force);
698 f.set_stopped();
699 assert_eq!(f.phase.get(), Phase::Stopped(manner));
700
701 let f = Flags::new(false);
703 f.begin_terminate(force);
704 f.enter_transport_shutdown();
705 f.set_stopped();
706 assert_eq!(f.phase.get(), Phase::Stopped(manner));
707 assert!(f.is_closed() && f.is_terminating());
708 assert_eq!(f.is_force_closing(), force);
709 }
710
711 let f = Flags::new(false);
713 f.enter_filters_stopping();
714 f.enter_transport_shutdown();
715 f.set_stopped();
716 assert_eq!(f.phase.get(), Phase::Stopped(Manner::Graceful));
717 assert!(!f.is_terminating());
718 }
719
720 #[test]
721 fn a_terminated_connection_is_left_alone() {
722 let f = Flags::new(false);
723 f.set_stopped();
724 assert!(!f.begin_terminate(true));
725 assert!(!f.is_force_closing());
726 assert!(!f.is_terminating());
727 }
728}