1use std::{
27 fmt::Debug,
28 sync::{
29 Arc, Mutex,
30 atomic::{AtomicU8, Ordering},
31 },
32};
33
34use crate::mode::ConnectionMode;
35
36#[derive(Clone)]
41pub struct SocketStateSink {
42 callback: Arc<dyn Fn(SocketState) + Send + Sync>,
43 transition_lock: Arc<Mutex<()>>,
44}
45
46impl SocketStateSink {
47 #[must_use]
53 pub fn new<F>(callback: F) -> Self
54 where
55 F: Fn(SocketState) + Send + Sync + 'static,
56 {
57 Self {
58 callback: Arc::new(callback),
59 transition_lock: Arc::new(Mutex::new(())),
60 }
61 }
62
63 pub(crate) fn transition(
64 &self,
65 value: &AtomicU8,
66 current: ConnectionMode,
67 next: ConnectionMode,
68 state: SocketState,
69 ) -> bool {
70 self.transition_result(value, current, next, state).is_ok()
71 }
72
73 pub(crate) fn transition_result(
74 &self,
75 value: &AtomicU8,
76 current: ConnectionMode,
77 next: ConnectionMode,
78 state: SocketState,
79 ) -> Result<(), ConnectionMode> {
80 let _guard = self
81 .transition_lock
82 .lock()
83 .expect("socket state sink transition lock poisoned");
84
85 if let Err(actual) = value.compare_exchange(
86 current.as_u8(),
87 next.as_u8(),
88 Ordering::SeqCst,
89 Ordering::SeqCst,
90 ) {
91 return Err(ConnectionMode::from_u8(actual));
92 }
93
94 self.notify(state);
95
96 Ok(())
97 }
98
99 pub(crate) fn publish_websocket(&self, state: SocketState) {
100 let _guard = self
101 .transition_lock
102 .lock()
103 .expect("socket state sink transition lock poisoned");
104 self.notify(state);
105 }
106
107 pub(crate) fn close_on_loss(&self, value: &AtomicU8) -> bool {
108 let _guard = self
109 .transition_lock
110 .lock()
111 .expect("socket state sink transition lock poisoned");
112 let current = ConnectionMode::from_atomic(value);
113
114 if !matches!(current, ConnectionMode::Active | ConnectionMode::Reconnect)
115 || value
116 .compare_exchange(
117 current.as_u8(),
118 ConnectionMode::Closed.as_u8(),
119 Ordering::SeqCst,
120 Ordering::SeqCst,
121 )
122 .is_err()
123 {
124 return false;
125 }
126
127 if current.is_active() {
128 self.notify(SocketState::Disconnected);
129 }
130
131 true
132 }
133
134 fn notify(&self, state: SocketState) {
135 if std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| (self.callback)(state)))
136 .is_err()
137 {
138 log::error!("Socket state sink panicked while handling {state:?}");
139 }
140 }
141}
142
143impl Debug for SocketStateSink {
144 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
145 f.debug_struct(stringify!(SocketStateSink))
146 .finish_non_exhaustive()
147 }
148}
149
150#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
152pub enum SocketState {
153 Connected,
155 Disconnected,
157}
158
159#[cfg(test)]
160mod tests {
161 use std::sync::{
162 Barrier, Mutex,
163 atomic::{AtomicU8, AtomicUsize, Ordering as AtomicOrdering},
164 };
165
166 use rstest::rstest;
167
168 use super::*;
169 use crate::mode::{ConnectionMode, ReconnectOutcome};
170
171 #[rstest]
172 fn state_sink_reports_only_successful_edges_in_order() {
173 let states = Arc::new(Mutex::new(Vec::new()));
174 let states_callback = Arc::clone(&states);
175 let sink = SocketStateSink::new(move |state| {
176 states_callback.lock().unwrap().push(state);
177 });
178 let mode = AtomicU8::new(ConnectionMode::Reconnect.as_u8());
179
180 assert_eq!(
181 ConnectionMode::complete_reconnect_with_sink(&mode, Some(&sink)),
182 ReconnectOutcome::Reconnected
183 );
184 assert!(ConnectionMode::request_reconnect_with_sink(
185 &mode,
186 Some(&sink)
187 ));
188 assert!(!ConnectionMode::request_reconnect_with_sink(
189 &mode,
190 Some(&sink)
191 ));
192 assert_eq!(
193 ConnectionMode::complete_reconnect_with_sink(&mode, Some(&sink)),
194 ReconnectOutcome::Reconnected
195 );
196
197 assert_eq!(
198 *states.lock().unwrap(),
199 vec![
200 SocketState::Connected,
201 SocketState::Disconnected,
202 SocketState::Connected,
203 ]
204 );
205 }
206
207 #[rstest]
208 fn state_sink_reports_one_concurrent_loss() {
209 let states = Arc::new(Mutex::new(Vec::new()));
210 let states_callback = Arc::clone(&states);
211 let sink = SocketStateSink::new(move |state| {
212 states_callback.lock().unwrap().push(state);
213 });
214 let mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
215 let barrier = Arc::new(Barrier::new(8));
216
217 let mut transitions = Vec::with_capacity(8);
218
219 for _ in 0..8 {
220 let mode = Arc::clone(&mode);
221 let sink = sink.clone();
222 let barrier = Arc::clone(&barrier);
223 transitions.push(std::thread::spawn(move || {
224 barrier.wait();
225 ConnectionMode::request_reconnect_with_sink(&mode, Some(&sink))
226 }));
227 }
228
229 let successful = transitions
230 .into_iter()
231 .map(|transition| transition.join().unwrap())
232 .filter(|successful| *successful)
233 .count();
234
235 assert_eq!(successful, 1);
236 assert_eq!(
237 ConnectionMode::from_atomic(&mode),
238 ConnectionMode::Reconnect
239 );
240 assert_eq!(*states.lock().unwrap(), vec![SocketState::Disconnected]);
241 }
242
243 #[rstest]
244 fn state_sink_reports_one_mixed_concurrent_loss() {
245 let states = Arc::new(Mutex::new(Vec::new()));
246 let states_callback = Arc::clone(&states);
247 let sink = SocketStateSink::new(move |state| {
248 states_callback.lock().unwrap().push(state);
249 });
250 let mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
251 let barrier = Arc::new(Barrier::new(2));
252
253 let reconnect = {
254 let mode = Arc::clone(&mode);
255 let sink = sink.clone();
256 let barrier = Arc::clone(&barrier);
257 std::thread::spawn(move || {
258 barrier.wait();
259 ConnectionMode::request_reconnect_with_sink(&mode, Some(&sink))
260 })
261 };
262
263 let close = std::thread::spawn({
264 let mode = Arc::clone(&mode);
265 move || {
266 barrier.wait();
267 ConnectionMode::close_websocket_on_loss(&mode, Some(&sink))
268 }
269 });
270
271 reconnect.join().unwrap();
272 let closed = close.join().unwrap();
273
274 assert!(closed);
275 assert_eq!(ConnectionMode::from_atomic(&mode), ConnectionMode::Closed);
276 assert_eq!(*states.lock().unwrap(), vec![SocketState::Disconnected]);
277 }
278
279 #[rstest]
280 fn state_sink_continues_after_callback_panic() {
281 let calls = Arc::new(AtomicUsize::new(0));
282 let calls_callback = Arc::clone(&calls);
283 let states = Arc::new(Mutex::new(Vec::new()));
284 let states_callback = Arc::clone(&states);
285 let sink = SocketStateSink::new(move |state| {
286 assert_ne!(
287 calls_callback.fetch_add(1, AtomicOrdering::SeqCst),
288 0,
289 "test socket state callback panic"
290 );
291 states_callback.lock().unwrap().push(state);
292 });
293 let mode = AtomicU8::new(ConnectionMode::Reconnect.as_u8());
294
295 assert_eq!(
296 ConnectionMode::complete_reconnect_with_sink(&mode, Some(&sink)),
297 ReconnectOutcome::Reconnected
298 );
299 assert!(ConnectionMode::request_reconnect_with_sink(
300 &mode,
301 Some(&sink)
302 ));
303
304 assert_eq!(calls.load(AtomicOrdering::SeqCst), 2);
305 assert_eq!(
306 ConnectionMode::from_atomic(&mode),
307 ConnectionMode::Reconnect
308 );
309 assert_eq!(*states.lock().unwrap(), vec![SocketState::Disconnected]);
310 }
311
312 #[rstest]
313 fn state_sink_suppresses_deliberate_disconnect() {
314 let states = Arc::new(Mutex::new(Vec::new()));
315 let states_callback = Arc::clone(&states);
316 let sink = SocketStateSink::new(move |state| {
317 states_callback.lock().unwrap().push(state);
318 });
319 let mode = AtomicU8::new(ConnectionMode::Active.as_u8());
320
321 assert!(ConnectionMode::request_disconnect(&mode));
322 assert!(!ConnectionMode::request_reconnect_with_sink(
323 &mode,
324 Some(&sink)
325 ));
326
327 assert_eq!(*states.lock().unwrap(), Vec::new());
328 }
329
330 #[rstest]
331 fn state_sink_closes_after_reported_loss_without_another_event() {
332 let states = Arc::new(Mutex::new(Vec::new()));
333 let states_callback = Arc::clone(&states);
334 let sink = SocketStateSink::new(move |state| {
335 states_callback.lock().unwrap().push(state);
336 });
337 let mode = AtomicU8::new(ConnectionMode::Reconnect.as_u8());
338
339 assert!(ConnectionMode::close_websocket_on_loss(&mode, Some(&sink)));
340
341 assert_eq!(ConnectionMode::from_atomic(&mode), ConnectionMode::Closed);
342 assert_eq!(*states.lock().unwrap(), Vec::new());
343 }
344}