Skip to main content

ant_quic/
fair_polling.rs

1// Copyright 2024 Saorsa Labs Ltd.
2//
3// This Saorsa Network Software is licensed under the General Public License (GPL), version 3.
4// Please see the file LICENSE-GPL, or visit <http://www.gnu.org/licenses/> for the full text.
5//
6// Full details available at https://saorsalabs.com/licenses
7
8//! Fair polling for multiple transports
9//!
10//! Prevents starvation by alternating poll order between
11//! direct and relay transports.
12
13use std::sync::atomic::{AtomicU64, Ordering};
14
15/// Order in which to poll transports
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub enum PollOrder {
18    /// Poll direct transports first, then relay
19    DirectFirst,
20    /// Poll relay transports first, then direct
21    RelayFirst,
22}
23
24/// Fair poller that alternates poll order to prevent starvation
25#[derive(Debug)]
26pub struct FairPoller {
27    counter: AtomicU64,
28}
29
30impl FairPoller {
31    /// Create a new fair poller
32    pub fn new() -> Self {
33        Self {
34            counter: AtomicU64::new(0),
35        }
36    }
37
38    /// Get the poll order for this iteration
39    ///
40    /// Increments counter and returns appropriate order.
41    /// Alternates between DirectFirst and RelayFirst to ensure
42    /// fair access to both transport types.
43    pub fn poll_order(&self) -> PollOrder {
44        let count = self.counter.fetch_add(1, Ordering::Relaxed);
45        if count % 2 == 0 {
46            PollOrder::DirectFirst
47        } else {
48            PollOrder::RelayFirst
49        }
50    }
51
52    /// Get the poll order without incrementing the counter
53    pub fn peek_order(&self) -> PollOrder {
54        let count = self.counter.load(Ordering::Relaxed);
55        if count % 2 == 0 {
56            PollOrder::DirectFirst
57        } else {
58            PollOrder::RelayFirst
59        }
60    }
61
62    /// Reset the counter
63    pub fn reset(&self) {
64        self.counter.store(0, Ordering::Relaxed);
65    }
66
67    /// Get the current counter value
68    pub fn counter(&self) -> u64 {
69        self.counter.load(Ordering::Relaxed)
70    }
71
72    /// Set counter (for testing)
73    #[cfg(test)]
74    pub fn set_counter(&self, value: u64) {
75        self.counter.store(value, Ordering::Relaxed);
76    }
77}
78
79impl Default for FairPoller {
80    fn default() -> Self {
81        Self::new()
82    }
83}
84
85/// Macro to poll transports in fair order
86///
87/// Usage:
88/// ```ignore
89/// poll_transports_fair!(
90///     poller,
91///     poll_direct_transport(),
92///     poll_relay_transport()
93/// )
94/// ```
95#[macro_export]
96macro_rules! poll_transports_fair {
97    ($poller:expr, $direct:expr, $relay:expr) => {{
98        use $crate::fair_polling::PollOrder;
99        match $poller.poll_order() {
100            PollOrder::DirectFirst => {
101                if let Some(result) = $direct {
102                    Some(result)
103                } else {
104                    $relay
105                }
106            }
107            PollOrder::RelayFirst => {
108                if let Some(result) = $relay {
109                    Some(result)
110                } else {
111                    $direct
112                }
113            }
114        }
115    }};
116}
117
118#[cfg(test)]
119mod tests {
120    use super::*;
121
122    #[test]
123    fn test_alternating_poll_order() {
124        let poller = FairPoller::new();
125
126        let order1 = poller.poll_order();
127        let order2 = poller.poll_order();
128        let order3 = poller.poll_order();
129        let order4 = poller.poll_order();
130
131        // Should alternate
132        assert_eq!(order1, PollOrder::DirectFirst);
133        assert_eq!(order2, PollOrder::RelayFirst);
134        assert_eq!(order3, PollOrder::DirectFirst);
135        assert_eq!(order4, PollOrder::RelayFirst);
136    }
137
138    #[test]
139    fn test_counter_wraps() {
140        let poller = FairPoller::new();
141
142        // Set counter near max
143        poller.set_counter(u64::MAX);
144
145        // Should wrap without panic
146        let _ = poller.poll_order();
147        let _ = poller.poll_order();
148
149        // Counter should have wrapped
150        assert!(poller.counter() < u64::MAX);
151    }
152
153    #[test]
154    fn test_poll_order_is_deterministic() {
155        let poller = FairPoller::new();
156
157        // Even counter: direct first
158        poller.set_counter(0);
159        assert_eq!(poller.poll_order(), PollOrder::DirectFirst);
160
161        // Reset and check odd
162        poller.set_counter(1);
163        assert_eq!(poller.poll_order(), PollOrder::RelayFirst);
164    }
165
166    #[test]
167    fn test_peek_does_not_increment() {
168        let poller = FairPoller::new();
169
170        let peek1 = poller.peek_order();
171        let peek2 = poller.peek_order();
172        let peek3 = poller.peek_order();
173
174        // Should all be the same since counter isn't incremented
175        assert_eq!(peek1, peek2);
176        assert_eq!(peek2, peek3);
177        assert_eq!(poller.counter(), 0);
178    }
179
180    #[test]
181    fn test_reset() {
182        let poller = FairPoller::new();
183
184        poller.poll_order();
185        poller.poll_order();
186        poller.poll_order();
187
188        assert_eq!(poller.counter(), 3);
189
190        poller.reset();
191        assert_eq!(poller.counter(), 0);
192        assert_eq!(poller.peek_order(), PollOrder::DirectFirst);
193    }
194
195    #[test]
196    fn test_default() {
197        let poller = FairPoller::default();
198        assert_eq!(poller.counter(), 0);
199    }
200
201    #[test]
202    fn test_poll_transports_fair_macro_direct_first() {
203        let poller = FairPoller::new();
204        poller.set_counter(0); // DirectFirst
205
206        let direct = Some(1);
207        let relay: Option<i32> = Some(2);
208
209        let result = poll_transports_fair!(poller, direct, relay);
210        assert_eq!(result, Some(1)); // Direct should be selected
211    }
212
213    #[test]
214    fn test_poll_transports_fair_macro_relay_first() {
215        let poller = FairPoller::new();
216        poller.set_counter(1); // RelayFirst
217
218        let direct: Option<i32> = Some(1);
219        let relay = Some(2);
220
221        let result = poll_transports_fair!(poller, direct, relay);
222        assert_eq!(result, Some(2)); // Relay should be selected
223    }
224
225    #[test]
226    fn test_poll_transports_fair_macro_fallback() {
227        let poller = FairPoller::new();
228        poller.set_counter(0); // DirectFirst
229
230        let direct: Option<i32> = None;
231        let relay = Some(2);
232
233        let result = poll_transports_fair!(poller, direct, relay);
234        assert_eq!(result, Some(2)); // Should fall back to relay
235    }
236
237    #[test]
238    fn test_poll_transports_fair_macro_both_none() {
239        let poller = FairPoller::new();
240
241        let direct: Option<i32> = None;
242        let relay: Option<i32> = None;
243
244        let result = poll_transports_fair!(poller, direct, relay);
245        assert_eq!(result, None);
246    }
247
248    #[test]
249    fn test_concurrent_access() {
250        use std::sync::Arc;
251        use std::thread;
252
253        let poller = Arc::new(FairPoller::new());
254        let mut handles = vec![];
255
256        for _ in 0..10 {
257            let p = Arc::clone(&poller);
258            handles.push(thread::spawn(move || {
259                for _ in 0..100 {
260                    let _ = p.poll_order();
261                }
262            }));
263        }
264
265        for handle in handles {
266            handle.join().expect("Thread panicked");
267        }
268
269        // Should have incremented 1000 times
270        assert_eq!(poller.counter(), 1000);
271    }
272}