Skip to main content

hopr_transport_p2p/
liveness.rs

1//! Per-peer connection liveness tracking for libp2p streams.
2//!
3//! [`LivenessStream`] wraps any `AsyncRead + AsyncWrite` value and checks an
4//! `Arc<AtomicBool>` liveness flag **before every poll**. When the swarm event
5//! loop clears the flag on `ConnectionClosed` / `OutgoingConnectionError`, the
6//! very next read or write on any wrapped substream for that peer returns
7//! `Err(io::ErrorKind::ConnectionAborted)`.
8//!
9//! This makes dead streams self-signal to their consumers (the per-peer reader
10//! and writer tasks in `hopr-transport`) without any changes to the
11//! `NetworkStreamControl` trait or the protocol layer.
12
13use std::{
14    io,
15    pin::Pin,
16    sync::{
17        Arc,
18        atomic::{AtomicBool, Ordering},
19    },
20    task::{Context, Poll},
21};
22
23use dashmap::DashMap;
24use futures::{AsyncRead, AsyncWrite};
25use libp2p::PeerId;
26use pin_project::pin_project;
27
28/// Shared registry mapping each peer to its connection-liveness flag.
29///
30/// A clone of this registry is held by both [`crate::HoprNetwork`] and the
31/// swarm event-loop closure. When libp2p reports a peer fully disconnected,
32/// the loop calls [`LivenessRegistry::remove`], which clears the flag and
33/// removes the entry. Any [`LivenessStream`] that already cloned the old
34/// `Arc<AtomicBool>` will see `false` on its next poll.
35#[derive(Clone, Default)]
36pub(crate) struct LivenessRegistry(Arc<DashMap<PeerId, Arc<AtomicBool>>>);
37
38impl LivenessRegistry {
39    /// Registers `peer` and returns its liveness flag, creating a fresh `true`
40    /// flag if none exists. Call this when opening or accepting a stream.
41    pub(crate) fn get_or_create_connected(&self, peer: &PeerId) -> Arc<AtomicBool> {
42        self.0
43            .entry(*peer)
44            .or_insert_with(|| Arc::new(AtomicBool::new(true)))
45            .clone()
46    }
47
48    /// Returns the current liveness flag for `peer`, or `None` if the peer is
49    /// not registered. Does not create an entry.
50    #[allow(dead_code)]
51    pub(crate) fn get(&self, peer: &PeerId) -> Option<Arc<AtomicBool>> {
52        self.0.get(peer).map(|r| r.clone())
53    }
54
55    /// Removes `peer` from the registry and sets its flag to `false`.
56    ///
57    /// Any [`LivenessStream`] still holding the old `Arc` will error on its
58    /// next poll. A subsequent [`Self::get_or_create_connected`] call mints a fresh `true` flag.
59    pub(crate) fn remove(&self, peer: &PeerId) {
60        if let Some((_, flag)) = self.0.remove(peer) {
61            flag.store(false, Ordering::Release);
62        }
63    }
64}
65
66/// An `AsyncRead + AsyncWrite` wrapper that errors once its connection-liveness flag is cleared.
67///
68/// Every `poll_read`, `poll_write`, `poll_flush`, and `poll_close` performs a
69/// single `Relaxed` load of the flag. If the flag is `false`, the poll returns
70/// `Err(io::ErrorKind::ConnectionAborted)` immediately, without touching the
71/// inner stream. This surfaces naturally to `FramedRead`/`FramedWrite` above it,
72/// ending the per-peer reader/writer forward-futures and triggering cache invalidation.
73#[pin_project]
74pub(crate) struct LivenessStream<S> {
75    #[pin]
76    inner: S,
77    alive: Arc<AtomicBool>,
78}
79
80impl<S> LivenessStream<S> {
81    pub(crate) fn new(inner: S, alive: Arc<AtomicBool>) -> Self {
82        Self { inner, alive }
83    }
84
85    fn dead_error() -> io::Error {
86        io::Error::new(io::ErrorKind::ConnectionAborted, "connection to peer was closed")
87    }
88}
89
90impl<S: AsyncRead> AsyncRead for LivenessStream<S> {
91    fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<usize>> {
92        let this = self.project();
93        if !this.alive.load(Ordering::Acquire) {
94            return Poll::Ready(Err(Self::dead_error()));
95        }
96        this.inner.poll_read(cx, buf)
97    }
98}
99
100impl<S: AsyncWrite> AsyncWrite for LivenessStream<S> {
101    fn poll_write(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
102        let this = self.project();
103        if !this.alive.load(Ordering::Acquire) {
104            return Poll::Ready(Err(Self::dead_error()));
105        }
106        this.inner.poll_write(cx, buf)
107    }
108
109    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
110        let this = self.project();
111        if !this.alive.load(Ordering::Acquire) {
112            return Poll::Ready(Err(Self::dead_error()));
113        }
114        this.inner.poll_flush(cx)
115    }
116
117    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
118        let this = self.project();
119        if !this.alive.load(Ordering::Acquire) {
120            return Poll::Ready(Err(Self::dead_error()));
121        }
122        this.inner.poll_close(cx)
123    }
124}
125
126#[cfg(test)]
127mod tests {
128    use std::sync::atomic::AtomicBool;
129
130    use anyhow::Context;
131    use async_channel_io::pipe;
132    use futures::{AsyncReadExt, AsyncWriteExt};
133
134    use super::*;
135
136    fn in_memory_pipe() -> (impl AsyncWrite, impl AsyncRead) {
137        let (writer, reader) = pipe();
138        (writer, reader)
139    }
140
141    // ---------------------------------------------------------------------------
142    // LivenessStream unit tests
143    // ---------------------------------------------------------------------------
144
145    #[tokio::test]
146    async fn liveness_stream_should_pass_through_reads_when_alive() -> anyhow::Result<()> {
147        let (mut raw_write, raw_read) = in_memory_pipe();
148        let alive = Arc::new(AtomicBool::new(true));
149        let mut stream = LivenessStream::new(raw_read, alive.clone());
150
151        raw_write.write_all(b"hello").await.context("write failed")?;
152        drop(raw_write);
153
154        let mut buf = vec![0u8; 5];
155        stream.read_exact(&mut buf).await.context("read failed")?;
156        assert_eq!(&buf, b"hello");
157        Ok(())
158    }
159
160    #[tokio::test]
161    async fn liveness_stream_should_pass_through_writes_when_alive() -> anyhow::Result<()> {
162        let (raw_write, mut raw_read) = in_memory_pipe();
163        let alive = Arc::new(AtomicBool::new(true));
164        let mut stream = LivenessStream::new(raw_write, alive.clone());
165
166        stream.write_all(b"world").await.context("write failed")?;
167        stream.flush().await.context("flush failed")?;
168        drop(stream);
169
170        let mut buf = vec![0u8; 5];
171        raw_read.read_exact(&mut buf).await.context("read failed")?;
172        assert_eq!(&buf, b"world");
173        Ok(())
174    }
175
176    #[tokio::test]
177    async fn liveness_stream_should_error_on_read_when_flag_cleared() -> anyhow::Result<()> {
178        let (_raw_write, raw_read) = in_memory_pipe();
179        let alive = Arc::new(AtomicBool::new(true));
180        let mut stream = LivenessStream::new(raw_read, alive.clone());
181
182        alive.store(false, Ordering::Relaxed);
183
184        let mut buf = vec![0u8; 4];
185        let result = stream.read(&mut buf).await;
186
187        assert!(
188            matches!(result, Err(ref e) if e.kind() == io::ErrorKind::ConnectionAborted),
189            "expected ConnectionAborted, got {result:?}"
190        );
191        Ok(())
192    }
193
194    #[tokio::test]
195    async fn liveness_stream_should_error_on_write_when_flag_cleared() -> anyhow::Result<()> {
196        let (raw_write, _raw_read) = in_memory_pipe();
197        let alive = Arc::new(AtomicBool::new(true));
198        let mut stream = LivenessStream::new(raw_write, alive.clone());
199
200        alive.store(false, Ordering::Relaxed);
201
202        let result = stream.write(b"data").await;
203
204        assert!(
205            matches!(result, Err(ref e) if e.kind() == io::ErrorKind::ConnectionAborted),
206            "expected ConnectionAborted, got {result:?}"
207        );
208        Ok(())
209    }
210
211    #[tokio::test]
212    async fn liveness_stream_should_error_on_flush_when_flag_cleared() -> anyhow::Result<()> {
213        let (raw_write, _raw_read) = in_memory_pipe();
214        let alive = Arc::new(AtomicBool::new(true));
215        let mut stream = LivenessStream::new(raw_write, alive.clone());
216
217        alive.store(false, Ordering::Relaxed);
218
219        let result = stream.flush().await;
220
221        assert!(
222            matches!(result, Err(ref e) if e.kind() == io::ErrorKind::ConnectionAborted),
223            "expected ConnectionAborted, got {result:?}"
224        );
225        Ok(())
226    }
227
228    #[tokio::test]
229    async fn liveness_stream_should_error_on_close_when_flag_cleared() -> anyhow::Result<()> {
230        let (raw_write, _raw_read) = in_memory_pipe();
231        let alive = Arc::new(AtomicBool::new(true));
232        let mut stream = LivenessStream::new(raw_write, alive.clone());
233
234        alive.store(false, Ordering::Relaxed);
235
236        let result = stream.close().await;
237
238        assert!(
239            matches!(result, Err(ref e) if e.kind() == io::ErrorKind::ConnectionAborted),
240            "expected ConnectionAborted, got {result:?}"
241        );
242        Ok(())
243    }
244
245    // ---------------------------------------------------------------------------
246    // LivenessRegistry tests
247    // ---------------------------------------------------------------------------
248
249    #[test]
250    fn get_or_create_connected_should_create_alive_flag_for_new_peer() {
251        let registry = LivenessRegistry::default();
252        let peer = PeerId::random();
253
254        let flag = registry.get_or_create_connected(&peer);
255
256        assert!(flag.load(Ordering::Relaxed), "new flag must be alive");
257        assert!(registry.0.contains_key(&peer));
258    }
259
260    #[test]
261    fn get_or_create_connected_should_return_same_flag_for_same_peer() {
262        let registry = LivenessRegistry::default();
263        let peer = PeerId::random();
264
265        let flag1 = registry.get_or_create_connected(&peer);
266        let flag2 = registry.get_or_create_connected(&peer);
267
268        assert!(Arc::ptr_eq(&flag1, &flag2), "should return the same Arc");
269    }
270
271    #[test]
272    fn get_should_return_none_for_unregistered_peer() {
273        let registry = LivenessRegistry::default();
274        let peer = PeerId::random();
275
276        assert!(registry.get(&peer).is_none());
277    }
278
279    #[test]
280    fn get_should_return_flag_after_get_or_create_connected() {
281        let registry = LivenessRegistry::default();
282        let peer = PeerId::random();
283
284        let get_or_create_connecteded = registry.get_or_create_connected(&peer);
285        let got = registry
286            .get(&peer)
287            .expect("flag must exist after get_or_create_connected");
288
289        assert!(
290            Arc::ptr_eq(&get_or_create_connecteded, &got),
291            "get must return the same Arc as get_or_create_connected"
292        );
293    }
294
295    #[test]
296    fn get_should_return_none_after_remove() {
297        let registry = LivenessRegistry::default();
298        let peer = PeerId::random();
299
300        registry.get_or_create_connected(&peer);
301        registry.remove(&peer);
302
303        assert!(registry.get(&peer).is_none());
304    }
305
306    #[test]
307    fn remove_should_clear_flag_and_remove_entry() {
308        let registry = LivenessRegistry::default();
309        let peer = PeerId::random();
310
311        let flag = registry.get_or_create_connected(&peer);
312        assert!(flag.load(Ordering::Relaxed));
313
314        registry.remove(&peer);
315
316        assert!(!flag.load(Ordering::Relaxed), "flag must be cleared after remove");
317        assert!(!registry.0.contains_key(&peer));
318    }
319
320    #[test]
321    fn remove_should_be_safe_when_peer_not_in_registry() {
322        let registry = LivenessRegistry::default();
323        let peer = PeerId::random();
324
325        registry.remove(&peer);
326    }
327
328    #[test]
329    fn subsequent_get_or_create_connected_after_remove_should_return_fresh_alive_flag() {
330        let registry = LivenessRegistry::default();
331        let peer = PeerId::random();
332
333        let old_flag = registry.get_or_create_connected(&peer);
334        registry.remove(&peer);
335
336        let new_flag = registry.get_or_create_connected(&peer);
337
338        assert!(!old_flag.load(Ordering::Relaxed), "old flag must be dead");
339        assert!(new_flag.load(Ordering::Relaxed), "new flag must be alive");
340        assert!(!Arc::ptr_eq(&old_flag, &new_flag), "must be a distinct Arc");
341    }
342}