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, Mutex,
30        atomic::{AtomicU8, Ordering},
31    },
32};
33
34use crate::mode::ConnectionMode;
35
36/// Receives ordered semantic state changes from a single socket client.
37///
38/// Clients must route every transition into or out of [`ConnectionMode::Active`] through the
39/// sink-backed transition methods so each edge is reported once.
40#[derive(Clone)]
41pub struct SocketStateSink {
42    callback: Arc<dyn Fn(SocketState) + Send + Sync>,
43    transition_lock: Arc<Mutex<()>>,
44}
45
46impl SocketStateSink {
47    /// Creates a new [`SocketStateSink`] instance.
48    ///
49    /// The callback runs synchronously with each successful state transition and should return
50    /// promptly. It must not initiate another transition using the same sink because callbacks are
51    /// serialized under a non-reentrant lock.
52    #[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/// Represents the availability state reported by a socket transport.
151#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
152pub enum SocketState {
153    /// The transport is available.
154    Connected,
155    /// An active transport was lost.
156    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}