hopr_transport_p2p/
liveness.rs1use 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#[derive(Clone, Default)]
36pub(crate) struct LivenessRegistry(Arc<DashMap<PeerId, Arc<AtomicBool>>>);
37
38impl LivenessRegistry {
39 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 #[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 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#[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 #[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 #[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}