Skip to main content

nautilus_network/
sink.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Ordered availability-edge publication for socket transports.
17//!
18//! # Ordering
19//!
20//! A [`SocketStateSink`] reports transitions into and out of [`ConnectionMode::Active`] for one
21//! client, rather than every internal connection mode. Most paths perform the mode transition and
22//! callback under the same serialization lock, so concurrent loss and recovery attempts publish
23//! at most one ordered edge. WebSocket reconnect changes mode before publishing loss and uses the
24//! same lock for publication, while the controller prevents recovery until publication completes.
25
26use 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/// Receives ordered semantic state changes from a single socket client.
39///
40/// Clients must route every transition into or out of [`ConnectionMode::Active`] through the
41/// sink-backed transition methods so each edge is reported once.
42#[derive(Clone)]
43pub struct SocketStateSink {
44    callback: Arc<dyn Fn(SocketState) + Send + Sync>,
45    transition_lock: Arc<Mutex<()>>,
46}
47
48impl SocketStateSink {
49    /// Creates a new [`SocketStateSink`] instance.
50    ///
51    /// The callback runs synchronously with each successful state transition and should return
52    /// promptly. It must not initiate another transition using the same sink because callbacks are
53    /// serialized under a non-reentrant lock.
54    #[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    /// Returns a sink that invokes `callback` before forwarding each state to this sink.
66    #[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/// Represents the availability state reported by a socket transport.
164#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
165pub enum SocketState {
166    /// The transport is available.
167    Connected,
168    /// An active transport was lost.
169    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}