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