1use std::sync::atomic::{AtomicU64, Ordering};
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub enum PollOrder {
18 DirectFirst,
20 RelayFirst,
22}
23
24#[derive(Debug)]
26pub struct FairPoller {
27 counter: AtomicU64,
28}
29
30impl FairPoller {
31 pub fn new() -> Self {
33 Self {
34 counter: AtomicU64::new(0),
35 }
36 }
37
38 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 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 pub fn reset(&self) {
64 self.counter.store(0, Ordering::Relaxed);
65 }
66
67 pub fn counter(&self) -> u64 {
69 self.counter.load(Ordering::Relaxed)
70 }
71
72 #[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_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 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 poller.set_counter(u64::MAX);
144
145 let _ = poller.poll_order();
147 let _ = poller.poll_order();
148
149 assert!(poller.counter() < u64::MAX);
151 }
152
153 #[test]
154 fn test_poll_order_is_deterministic() {
155 let poller = FairPoller::new();
156
157 poller.set_counter(0);
159 assert_eq!(poller.poll_order(), PollOrder::DirectFirst);
160
161 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 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); 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)); }
212
213 #[test]
214 fn test_poll_transports_fair_macro_relay_first() {
215 let poller = FairPoller::new();
216 poller.set_counter(1); 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)); }
224
225 #[test]
226 fn test_poll_transports_fair_macro_fallback() {
227 let poller = FairPoller::new();
228 poller.set_counter(0); 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)); }
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 assert_eq!(poller.counter(), 1000);
271 }
272}