1use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
35use std::sync::{Arc, OnceLock, PoisonError};
36use std::time::Duration;
37
38use futures::future::BoxFuture;
39use tokio::sync::{Mutex, RwLock};
40use tokio::time::sleep;
41use tokio_util::sync::CancellationToken;
42use tracing::{info, warn};
43
44use aranet_types::{CurrentReading, DeviceInfo, DeviceType, HistoryRecord};
45
46use crate::connector::{ConnectFn, SensorLink, ble_connector, log_failed_release, release_link};
47use crate::device::Device;
48use crate::error::{Error, Result};
49use crate::events::{DeviceEvent, DeviceId, EventSender};
50use crate::history::{HistoryInfo, HistoryOptions};
51use crate::settings::{CalibrationData, MeasurementInterval};
52use crate::traits::AranetDevice;
53
54#[derive(Debug, Clone)]
56pub struct ReconnectOptions {
57 pub max_attempts: Option<u32>,
59 pub initial_delay: Duration,
61 pub max_delay: Duration,
63 pub backoff_multiplier: f64,
65 pub use_exponential_backoff: bool,
67}
68
69impl Default for ReconnectOptions {
70 fn default() -> Self {
71 Self {
72 max_attempts: Some(5),
73 initial_delay: Duration::from_secs(1),
74 max_delay: Duration::from_secs(60),
75 backoff_multiplier: 2.0,
76 use_exponential_backoff: true,
77 }
78 }
79}
80
81impl ReconnectOptions {
82 pub fn new() -> Self {
84 Self::default()
85 }
86
87 pub fn unlimited() -> Self {
89 Self {
90 max_attempts: None,
91 ..Default::default()
92 }
93 }
94
95 pub fn fixed_delay(delay: Duration) -> Self {
97 Self {
98 initial_delay: delay,
99 use_exponential_backoff: false,
100 ..Default::default()
101 }
102 }
103
104 pub fn max_attempts(mut self, attempts: u32) -> Self {
106 self.max_attempts = Some(attempts);
107 self
108 }
109
110 pub fn initial_delay(mut self, delay: Duration) -> Self {
112 self.initial_delay = delay;
113 self
114 }
115
116 pub fn max_delay(mut self, delay: Duration) -> Self {
118 self.max_delay = delay;
119 self
120 }
121
122 pub fn backoff_multiplier(mut self, multiplier: f64) -> Self {
124 self.backoff_multiplier = multiplier;
125 self
126 }
127
128 pub fn exponential_backoff(mut self, enabled: bool) -> Self {
130 self.use_exponential_backoff = enabled;
131 self
132 }
133
134 pub fn delay_for_attempt(&self, attempt: u32) -> Duration {
136 if !self.use_exponential_backoff {
137 return self.initial_delay;
138 }
139
140 let capped_attempt = attempt.min(32);
143 let delay_ms = self.initial_delay.as_millis() as f64
144 * self.backoff_multiplier.powi(capped_attempt as i32);
145
146 let delay = if delay_ms.is_finite() && delay_ms <= u64::MAX as f64 {
148 Duration::from_millis(delay_ms as u64)
149 } else {
150 self.max_delay
151 };
152
153 delay.min(self.max_delay)
154 }
155
156 pub fn validate(&self) -> Result<()> {
163 if self.backoff_multiplier < 1.0 {
164 return Err(Error::InvalidConfig(
165 "backoff_multiplier must be >= 1.0".to_string(),
166 ));
167 }
168 if self.initial_delay.is_zero() {
169 return Err(Error::InvalidConfig(
170 "initial_delay must be > 0".to_string(),
171 ));
172 }
173 if self.max_delay < self.initial_delay {
174 return Err(Error::InvalidConfig(
175 "max_delay must be >= initial_delay".to_string(),
176 ));
177 }
178 Ok(())
179 }
180}
181
182#[derive(Debug, Clone, Copy, PartialEq, Eq)]
184pub enum ConnectionState {
185 Connected,
187 Disconnected,
189 Reconnecting,
191 Failed,
193}
194
195pub(crate) struct ReconnectCore<L: SensorLink> {
228 identifier: String,
229 connect: ConnectFn<L>,
230 link: RwLock<Option<Arc<L>>>,
231 generation: AtomicU64,
234 disconnecting: AtomicU32,
236 serial: Arc<Mutex<()>>,
239 wake: std::sync::Mutex<CancellationToken>,
242 cancelled: AtomicBool,
244 state: RwLock<ConnectionState>,
245 attempt_count: AtomicU32,
246 options: ReconnectOptions,
247 events: Option<EventSender>,
248}
249
250impl<L: SensorLink> ReconnectCore<L> {
251 pub(crate) fn new(
252 identifier: impl Into<String>,
253 connect: ConnectFn<L>,
254 link: Arc<L>,
255 options: ReconnectOptions,
256 ) -> Self {
257 Self {
258 identifier: identifier.into(),
259 connect,
260 link: RwLock::new(Some(link)),
261 generation: AtomicU64::new(0),
262 disconnecting: AtomicU32::new(0),
263 serial: Arc::new(Mutex::new(())),
264 wake: std::sync::Mutex::new(CancellationToken::new()),
265 cancelled: AtomicBool::new(false),
266 state: RwLock::new(ConnectionState::Connected),
267 attempt_count: AtomicU32::new(0),
268 options,
269 events: None,
270 }
271 }
272
273 pub(crate) fn set_events(&mut self, events: EventSender) {
274 self.events = Some(events);
275 }
276
277 pub(crate) fn identifier(&self) -> &str {
278 &self.identifier
279 }
280
281 fn lock_wake(&self) -> std::sync::MutexGuard<'_, CancellationToken> {
283 self.wake.lock().unwrap_or_else(PoisonError::into_inner)
284 }
285
286 pub(crate) fn cancel(&self) {
287 self.cancelled.store(true, Ordering::SeqCst);
288 self.lock_wake().cancel();
289 }
290
291 pub(crate) fn is_cancelled(&self) -> bool {
292 self.cancelled.load(Ordering::SeqCst)
293 }
294
295 pub(crate) fn reset_cancellation(&self) {
296 self.cancelled.store(false, Ordering::SeqCst);
297 }
298
299 pub(crate) async fn state(&self) -> ConnectionState {
300 *self.state.read().await
301 }
302
303 pub(crate) fn attempt_count(&self) -> u32 {
304 self.attempt_count.load(Ordering::SeqCst)
305 }
306
307 pub(crate) async fn link(&self) -> Option<Arc<L>> {
308 self.link.read().await.clone()
309 }
310
311 pub(crate) async fn is_connected(&self) -> bool {
312 match self.link().await {
313 Some(link) => link.is_connected().await,
314 None => false,
315 }
316 }
317
318 pub(crate) async fn run<T, F>(&self, op: F) -> Result<T>
329 where
330 F: for<'b> Fn(&'b L) -> BoxFuture<'b, Result<T>> + Send + Sync,
331 T: Send,
332 {
333 let seen = {
334 let guard = self.link.read().await;
335 let seen = self.generation.load(Ordering::SeqCst);
336 if let Some(link) = guard.as_ref()
337 && link.is_connected().await
338 {
339 match op(link).await {
340 Ok(value) => return Ok(value),
341 Err(e) if !e.is_connection_error() => return Err(e),
342 Err(e) => warn!("Operation failed with a connection error: {e}; reconnecting"),
343 }
344 }
345 seen
346 }; self.recover(seen).await?;
349
350 let guard = self.link.read().await;
351 match guard.as_ref() {
352 Some(link) => op(link).await,
353 None => Err(self.stopped_error()),
354 }
355 }
356
357 pub(crate) async fn run_owned<T, F, Fut>(&self, op: F) -> Result<T>
360 where
361 F: Fn(&L) -> Fut,
362 Fut: Future<Output = Result<T>>,
363 {
364 let seen = {
365 let guard = self.link.read().await;
366 let seen = self.generation.load(Ordering::SeqCst);
367 if let Some(link) = guard.as_ref()
368 && link.is_connected().await
369 {
370 match op(link).await {
371 Ok(value) => return Ok(value),
372 Err(e) if !e.is_connection_error() => return Err(e),
373 Err(e) => warn!("Operation failed with a connection error: {e}; reconnecting"),
374 }
375 }
376 seen
377 };
378
379 self.recover(seen).await?;
380
381 let guard = self.link.read().await;
382 match guard.as_ref() {
383 Some(link) => op(link).await,
384 None => Err(self.stopped_error()),
385 }
386 }
387
388 pub(crate) async fn reconnect(&self) -> Result<()> {
392 self.recover(self.generation.load(Ordering::SeqCst)).await?;
393 if self.link.read().await.is_some() {
394 Ok(())
395 } else {
396 Err(self.stopped_error())
397 }
398 }
399
400 fn stopped_error(&self) -> Error {
404 if self.is_cancelled() {
405 Error::Cancelled
406 } else {
407 Error::NotConnected
408 }
409 }
410
411 pub(crate) async fn disconnect(&self) -> Result<()> {
412 self.generation.fetch_add(1, Ordering::SeqCst);
415 let _in_progress = DisconnectInProgress::start(&self.disconnecting);
416 self.lock_wake().cancel();
417
418 let serial = Arc::new(Arc::clone(&self.serial).lock_owned().await);
421 let link = {
422 let mut link = self.link.write().await;
423 self.generation.fetch_add(1, Ordering::SeqCst);
424 link.take()
425 };
426 *self.state.write().await = ConnectionState::Disconnected;
428 match link {
429 Some(link) => release_link(link, Arc::clone(&serial)).await,
430 None => Ok(()),
431 }
432 }
433
434 async fn recover(&self, seen: u64) -> Result<()> {
438 let serial = Arc::new(Arc::clone(&self.serial).lock_owned().await);
442 let wake = {
444 let mut wake = self.lock_wake();
445 *wake = CancellationToken::new();
446 wake.clone()
447 };
448 if self.disconnecting.load(Ordering::SeqCst) > 0 {
449 return Err(self.stopped_error());
451 }
452 if self.generation.load(Ordering::SeqCst) != seen {
453 return Ok(());
456 }
457 if self.is_cancelled() {
458 return self.cancelled_reconnect().await;
459 }
460
461 *self.state.write().await = ConnectionState::Reconnecting;
462 self.attempt_count.store(0, Ordering::SeqCst);
463
464 let old = self.link.write().await.take();
469 if let Some(old) = old
470 && let Err(e) = release_link(old, Arc::clone(&serial)).await
471 {
472 log_failed_release("Closing the old link", &self.identifier, &e);
473 }
474
475 loop {
476 if self.is_cancelled() || wake.is_cancelled() {
479 return self.cancelled_reconnect().await;
480 }
481
482 let attempt = self.attempt_count.fetch_add(1, Ordering::SeqCst) + 1;
483 if let Some(max) = self.options.max_attempts
484 && attempt > max
485 {
486 *self.state.write().await = ConnectionState::Failed;
487 self.generation.fetch_add(1, Ordering::SeqCst);
489 return Err(Error::Timeout {
490 operation: format!("reconnect to '{}'", self.identifier),
491 duration: self.options.max_delay.saturating_mul(max),
494 });
495 }
496
497 if let Some(sender) = &self.events {
498 let _ = sender.send(DeviceEvent::ReconnectStarted {
499 device: DeviceId::new(&self.identifier),
500 attempt,
501 });
502 }
503 info!("Reconnection attempt {} for {}", attempt, self.identifier);
504
505 let delay = self.options.delay_for_attempt(attempt - 1);
506 if wake.run_until_cancelled(sleep(delay)).await.is_none() {
507 return self.cancelled_reconnect().await;
508 }
509
510 let new = match wake
512 .run_until_cancelled((self.connect)(&self.identifier))
513 .await
514 {
515 None => return self.cancelled_reconnect().await,
516 Some(Err(e)) => {
517 warn!("Reconnection attempt {} failed: {}", attempt, e);
518 continue;
519 }
520 Some(Ok(new)) => Arc::new(new),
521 };
522
523 let mut link = self.link.write().await;
524 if self.generation.load(Ordering::SeqCst) != seen || wake.is_cancelled() {
525 drop(link);
527 if let Err(e) = release_link(new, Arc::clone(&serial)).await {
528 log_failed_release("Closing the new link", &self.identifier, &e);
529 }
530 return self.cancelled_reconnect().await;
531 }
532 *link = Some(new);
533 self.generation.fetch_add(1, Ordering::SeqCst);
534 drop(link);
535 *self.state.write().await = ConnectionState::Connected;
536
537 if let Some(sender) = &self.events {
538 let _ = sender.send(DeviceEvent::ReconnectSucceeded {
539 device: DeviceId::new(&self.identifier),
540 attempts: attempt,
541 });
542 }
543 info!("Reconnected successfully after {} attempts", attempt);
544 return Ok(());
545 }
546 }
547
548 async fn cancelled_reconnect(&self) -> Result<()> {
549 *self.state.write().await = ConnectionState::Disconnected;
550 info!("Reconnection cancelled for {}", self.identifier);
551 Err(Error::Cancelled)
552 }
553}
554
555struct DisconnectInProgress<'a>(&'a AtomicU32);
558
559impl<'a> DisconnectInProgress<'a> {
560 fn start(count: &'a AtomicU32) -> Self {
561 count.fetch_add(1, Ordering::SeqCst);
562 Self(count)
563 }
564}
565
566impl Drop for DisconnectInProgress<'_> {
567 fn drop(&mut self) {
568 self.0.fetch_sub(1, Ordering::SeqCst);
569 }
570}
571
572pub struct ReconnectingDevice {
578 core: ReconnectCore<Device>,
579 cached_name: OnceLock<String>,
581 cached_device_type: OnceLock<DeviceType>,
583}
584
585impl ReconnectingDevice {
586 pub async fn connect(identifier: &str, options: ReconnectOptions) -> Result<Self> {
594 options.validate()?;
595 let connect = ble_connector();
596 let device = Arc::new(connect(identifier).await?);
597
598 let cached_name = OnceLock::new();
600 if let Some(name) = device.name() {
601 let _ = cached_name.set(name.to_string());
602 }
603
604 let cached_device_type = OnceLock::new();
605 if let Some(device_type) = device.device_type() {
606 let _ = cached_device_type.set(device_type);
607 }
608
609 Ok(Self {
610 core: ReconnectCore::new(identifier, connect, device, options),
611 cached_name,
612 cached_device_type,
613 })
614 }
615
616 pub async fn connect_with_events(
620 identifier: &str,
621 options: ReconnectOptions,
622 event_sender: EventSender,
623 ) -> Result<Self> {
624 let mut this = Self::connect(identifier, options).await?;
625 this.core.set_events(event_sender);
626 Ok(this)
627 }
628
629 pub fn cancel_reconnect(&self) {
639 self.core.cancel();
640 }
641
642 pub fn is_cancelled(&self) -> bool {
644 self.core.is_cancelled()
645 }
646
647 pub fn reset_cancellation(&self) {
654 self.core.reset_cancellation();
655 }
656
657 pub async fn state(&self) -> ConnectionState {
659 self.core.state().await
660 }
661
662 pub async fn is_connected(&self) -> bool {
664 self.core.is_connected().await
665 }
666
667 pub fn identifier(&self) -> &str {
669 self.core.identifier()
670 }
671
672 pub async fn with_device<F, Fut, T>(&self, f: F) -> Result<T>
687 where
688 F: Fn(&Device) -> Fut,
689 Fut: Future<Output = Result<T>>,
690 {
691 self.core.run_owned(f).await
692 }
693
694 pub async fn reconnect(&self) -> Result<()> {
718 self.core.reconnect().await
719 }
720
721 pub async fn disconnect(&self) -> Result<()> {
733 self.core.disconnect().await
734 }
735
736 pub async fn attempt_count(&self) -> u32 {
738 self.core.attempt_count()
739 }
740
741 pub async fn name(&self) -> Option<String> {
743 let device = self.core.link().await?;
744 device.name().map(str::to_string)
745 }
746
747 pub async fn address(&self) -> String {
749 match self.core.link().await {
750 Some(device) => device.address().to_string(),
751 None => self.core.identifier().to_string(),
752 }
753 }
754
755 pub async fn device_type(&self) -> Option<DeviceType> {
757 self.core.link().await?.device_type()
758 }
759}
760
761impl AranetDevice for ReconnectingDevice {
763 async fn is_connected(&self) -> bool {
764 ReconnectingDevice::is_connected(self).await
765 }
766
767 async fn connect(&self) -> Result<()> {
768 if self.is_connected().await {
770 return Ok(());
771 }
772 self.reconnect().await
774 }
775
776 async fn disconnect(&self) -> Result<()> {
777 ReconnectingDevice::disconnect(self).await
778 }
779
780 fn name(&self) -> Option<&str> {
781 self.cached_name.get().map(|s| s.as_str())
782 }
783
784 fn address(&self) -> &str {
785 self.core.identifier()
786 }
787
788 fn device_type(&self) -> Option<DeviceType> {
789 self.cached_device_type.get().copied()
790 }
791
792 async fn read_current(&self) -> Result<CurrentReading> {
793 self.core.run(|d| Box::pin(d.read_current())).await
794 }
795
796 async fn read_device_info(&self) -> Result<DeviceInfo> {
797 self.core.run(|d| Box::pin(d.read_device_info())).await
798 }
799
800 async fn read_rssi(&self) -> Result<i16> {
801 self.core.run(|d| Box::pin(d.read_rssi())).await
802 }
803
804 async fn read_battery(&self) -> Result<u8> {
805 self.core.run(|d| Box::pin(d.read_battery())).await
806 }
807
808 async fn get_history_info(&self) -> Result<HistoryInfo> {
809 self.core.run(|d| Box::pin(d.get_history_info())).await
810 }
811
812 async fn download_history(&self) -> Result<Vec<HistoryRecord>> {
813 self.core.run(|d| Box::pin(d.download_history())).await
814 }
815
816 async fn download_history_with_options(
817 &self,
818 options: HistoryOptions,
819 ) -> Result<Vec<HistoryRecord>> {
820 let opts = options.clone();
821 self.core
822 .run(move |d| {
823 let opts = opts.clone();
824 Box::pin(async move { d.download_history_with_options(opts).await })
825 })
826 .await
827 }
828
829 async fn get_interval(&self) -> Result<MeasurementInterval> {
830 self.core.run(|d| Box::pin(d.get_interval())).await
831 }
832
833 async fn set_interval(&self, interval: MeasurementInterval) -> Result<()> {
834 self.core
835 .run(move |d| Box::pin(d.set_interval(interval)))
836 .await
837 }
838
839 async fn get_calibration(&self) -> Result<CalibrationData> {
840 self.core.run(|d| Box::pin(d.get_calibration())).await
841 }
842}
843
844#[cfg(test)]
845mod tests {
846 use super::*;
847
848 #[test]
849 fn test_reconnect_options_default() {
850 let opts = ReconnectOptions::default();
851 assert_eq!(opts.max_attempts, Some(5));
852 assert!(opts.use_exponential_backoff);
853 }
854
855 #[test]
856 fn test_reconnect_options_unlimited() {
857 let opts = ReconnectOptions::unlimited();
858 assert!(opts.max_attempts.is_none());
859 }
860
861 #[test]
862 fn test_delay_calculation() {
863 let opts = ReconnectOptions {
864 initial_delay: Duration::from_secs(1),
865 max_delay: Duration::from_secs(60),
866 backoff_multiplier: 2.0,
867 use_exponential_backoff: true,
868 ..Default::default()
869 };
870
871 assert_eq!(opts.delay_for_attempt(0), Duration::from_secs(1));
872 assert_eq!(opts.delay_for_attempt(1), Duration::from_secs(2));
873 assert_eq!(opts.delay_for_attempt(2), Duration::from_secs(4));
874 assert_eq!(opts.delay_for_attempt(3), Duration::from_secs(8));
875 }
876
877 #[test]
878 fn test_delay_capped_at_max() {
879 let opts = ReconnectOptions {
880 initial_delay: Duration::from_secs(1),
881 max_delay: Duration::from_secs(10),
882 backoff_multiplier: 2.0,
883 use_exponential_backoff: true,
884 ..Default::default()
885 };
886
887 assert_eq!(opts.delay_for_attempt(10), Duration::from_secs(10));
889 }
890
891 #[test]
892 fn test_fixed_delay() {
893 let opts = ReconnectOptions::fixed_delay(Duration::from_secs(5));
894 assert_eq!(opts.delay_for_attempt(0), Duration::from_secs(5));
895 assert_eq!(opts.delay_for_attempt(5), Duration::from_secs(5));
896 }
897}
898
899#[cfg(test)]
900mod lifecycle_tests {
901 use tokio::time::Instant;
902
903 use super::*;
904 use crate::test_support::{FakeConn, FakeEvent, FakeRadio, within};
905
906 const LIMIT: Duration = Duration::from_secs(600);
907
908 fn run_op(link: &FakeConn) -> BoxFuture<'_, Result<()>> {
910 Box::pin(link.op())
911 }
912
913 fn position(events: &[(Duration, FakeEvent)], wanted: &FakeEvent) -> Option<usize> {
914 events.iter().position(|(_, event)| event == wanted)
915 }
916
917 fn spawn_run(core: &Arc<ReconnectCore<FakeConn>>) -> tokio::task::JoinHandle<Result<()>> {
918 let core = Arc::clone(core);
919 tokio::spawn(async move { core.run(run_op).await })
920 }
921
922 async fn connected_core(
924 radio: &FakeRadio,
925 options: ReconnectOptions,
926 ) -> ReconnectCore<FakeConn> {
927 let connect = radio.connector();
928 let first = connect("A").await.expect("first connect");
929 ReconnectCore::new("A", connect, Arc::new(first), options)
930 }
931
932 #[tokio::test(start_paused = true)]
933 async fn reconnect_disconnects_the_old_link_before_connecting() {
934 within(LIMIT, async {
935 let radio = FakeRadio::new();
936 let core = connected_core(&radio, ReconnectOptions::default()).await;
937 radio.lose_link("A");
938
939 let result = core.run(run_op).await;
940
941 let events = radio.events();
942 assert!(result.is_ok(), "{result:?} after {events:#?}");
943 let closed = position(
944 &events,
945 &FakeEvent::Disconnect {
946 id: "A".into(),
947 handle: 1,
948 },
949 );
950 let opened = position(
951 &events,
952 &FakeEvent::Connected {
953 id: "A".into(),
954 handle: 2,
955 },
956 );
957 assert!(
958 matches!((closed, opened), (Some(closed), Some(opened)) if closed < opened),
959 "handle 1 must be disconnected before handle 2 connects: {events:#?}"
960 );
961 assert_eq!(core.link().await.map(|link| link.handle()), Some(2));
962 radio.assert_no_drop_teardown();
963 })
964 .await;
965 }
966
967 #[tokio::test(start_paused = true)]
968 async fn reconnected_link_survives_the_old_handle() {
969 within(LIMIT, async {
970 let radio = FakeRadio::new();
971 let core = connected_core(&radio, ReconnectOptions::default()).await;
972 radio.lose_link("A");
973
974 let first = core.run(run_op).await;
975 sleep(Duration::from_secs(1)).await;
977 let second = core.run(run_op).await;
978
979 assert_eq!(radio.connect_count("A"), 2, "{:#?}", radio.events());
980 assert!(first.is_ok() && second.is_ok(), "{first:?}, {second:?}");
981 })
982 .await;
983 }
984
985 #[tokio::test(start_paused = true)]
986 async fn non_connection_error_is_returned_without_reconnecting() {
987 within(LIMIT, async {
988 let radio = FakeRadio::new();
989 let core = connected_core(&radio, ReconnectOptions::default()).await;
990 let calls = AtomicU32::new(0);
991 let start = Instant::now();
992
993 let missing = core
994 .run(|_| {
995 calls.fetch_add(1, Ordering::SeqCst);
996 Box::pin(async {
997 Err::<(), _>(Error::CharacteristicNotFound {
998 uuid: "f0cd1502".into(),
999 service_count: 3,
1000 })
1001 })
1002 })
1003 .await;
1004 let invalid = core
1005 .run(|_| {
1006 calls.fetch_add(1, Ordering::SeqCst);
1007 Box::pin(async { Err::<(), _>(Error::InvalidData("bad".into())) })
1008 })
1009 .await;
1010
1011 assert!(
1012 matches!(&missing, Err(Error::CharacteristicNotFound { uuid, service_count: 3 }) if uuid == "f0cd1502"),
1013 "{missing:?}"
1014 );
1015 assert!(matches!(&invalid, Err(Error::InvalidData(msg)) if msg == "bad"), "{invalid:?}");
1016 assert_eq!(radio.connect_count("A"), 1, "{:#?}", radio.events());
1017 assert_eq!(calls.load(Ordering::SeqCst), 2);
1018 assert_eq!(start.elapsed(), Duration::ZERO);
1019 assert_eq!(core.link().await.map(|link| link.handle()), Some(1));
1020 })
1021 .await;
1022 }
1023
1024 #[tokio::test(start_paused = true)]
1025 async fn connection_error_reruns_the_operation_once() {
1026 within(LIMIT, async {
1027 let radio = FakeRadio::new();
1028 let core = connected_core(&radio, ReconnectOptions::default()).await;
1029 let calls = AtomicU32::new(0);
1030
1031 let result = core
1032 .run(|link| {
1033 if calls.fetch_add(1, Ordering::SeqCst) == 0 {
1034 Box::pin(async { Err(Error::NotConnected) })
1035 } else {
1036 Box::pin(link.op())
1037 }
1038 })
1039 .await;
1040
1041 assert!(result.is_ok(), "{result:?} after {:#?}", radio.events());
1042 assert_eq!(calls.load(Ordering::SeqCst), 2);
1043 assert_eq!(radio.connect_count("A"), 2);
1044 })
1045 .await;
1046 }
1047
1048 #[tokio::test(start_paused = true)]
1049 async fn concurrent_failures_share_one_reconnect() {
1050 within(LIMIT, async {
1051 let radio = FakeRadio::new();
1052 let core = connected_core(&radio, ReconnectOptions::default()).await;
1053 radio.lose_link("A");
1054
1055 let (first, second) = tokio::join!(core.run(run_op), core.run(run_op));
1056
1057 assert_eq!(radio.connect_count("A"), 2, "{:#?}", radio.events());
1058 assert!(first.is_ok() && second.is_ok(), "{first:?}, {second:?}");
1059 radio.assert_no_drop_teardown();
1060 })
1061 .await;
1062 }
1063
1064 #[tokio::test(start_paused = true)]
1065 async fn disconnect_during_backoff_stops_the_loop() {
1066 within(LIMIT, async {
1067 let radio = FakeRadio::new();
1068 let options = ReconnectOptions::default().initial_delay(Duration::from_secs(30));
1069 let core = Arc::new(connected_core(&radio, options).await);
1070 radio.script_connects("A", [false, false, true]);
1071 radio.lose_link("A");
1072 let start = Instant::now();
1073
1074 let recovering = spawn_run(&core);
1075 sleep(Duration::from_millis(500)).await;
1076 let waiting = spawn_run(&core);
1078 let reconnecting = tokio::spawn({
1079 let core = Arc::clone(&core);
1080 async move { core.reconnect().await }
1081 });
1082 sleep(Duration::from_millis(500)).await;
1083 let disconnected = core.disconnect().await;
1084
1085 let recovering = recovering.await.expect("join");
1086 let waiting = waiting.await.expect("join");
1087 let reconnecting = reconnecting.await.expect("join");
1088 assert!(
1089 matches!(recovering, Err(Error::Cancelled)),
1090 "{recovering:?}"
1091 );
1092 assert!(matches!(waiting, Err(Error::NotConnected)), "{waiting:?}");
1093 assert!(
1094 matches!(reconnecting, Err(Error::NotConnected)),
1095 "{reconnecting:?}"
1096 );
1097 assert!(disconnected.is_ok(), "{disconnected:?}");
1098 assert!(
1099 start.elapsed() < Duration::from_secs(30),
1100 "{:?}",
1101 start.elapsed()
1102 );
1103 assert_eq!(core.state().await, ConnectionState::Disconnected);
1104
1105 sleep(Duration::from_secs(300)).await;
1107 let events = radio.events();
1108 assert!(!radio.link_up("A"), "{events:#?}");
1109 assert!(
1110 !events.iter().any(
1111 |(_, event)| matches!(event, FakeEvent::Connected { handle, .. } if *handle > 1)
1112 ),
1113 "reconnected after disconnect(): {events:#?}"
1114 );
1115 assert_eq!(radio.connect_count("A"), 1, "{events:#?}");
1116 })
1117 .await;
1118 }
1119
1120 #[tokio::test(start_paused = true)]
1121 async fn cancel_reconnect_interrupts_the_backoff_sleep() {
1122 within(LIMIT, async {
1123 let radio = FakeRadio::new();
1124 let options = ReconnectOptions::default().initial_delay(Duration::from_secs(30));
1125 let core = Arc::new(connected_core(&radio, options).await);
1126
1127 let in_flight = tokio::spawn({
1130 let core = Arc::clone(&core);
1131 async move {
1132 core.run(|link| {
1133 Box::pin(async move {
1134 sleep(Duration::from_secs(1)).await;
1135 link.op().await
1136 })
1137 })
1138 .await
1139 }
1140 });
1141 sleep(Duration::from_millis(500)).await;
1142 radio.lose_link("A");
1143
1144 let start = Instant::now();
1145 let recovering = spawn_run(&core);
1146 sleep(Duration::from_secs(1)).await;
1147 core.cancel();
1148 let result = recovering.await.expect("join");
1149 assert!(matches!(result, Err(Error::Cancelled)), "{result:?}");
1150 assert!(
1151 start.elapsed() < Duration::from_secs(30),
1152 "cancelled after {:?}",
1153 start.elapsed()
1154 );
1155 let in_flight = in_flight.await.expect("join");
1156 assert!(matches!(in_flight, Err(Error::Cancelled)), "{in_flight:?}");
1157 assert_eq!(core.state().await, ConnectionState::Disconnected);
1158
1159 core.reset_cancellation();
1161 radio.set_connect_delay("A", Duration::from_secs(60));
1162 let start = Instant::now();
1163 let recovering = spawn_run(&core);
1164 sleep(Duration::from_secs(40)).await; core.cancel();
1166 let result = recovering.await.expect("join");
1167 assert!(matches!(result, Err(Error::Cancelled)), "{result:?}");
1168 assert_eq!(start.elapsed(), Duration::from_secs(40));
1169 assert_eq!(
1170 radio.connect_count("A"),
1171 2,
1172 "the first link and the cancelled connect"
1173 );
1174 sleep(Duration::from_secs(120)).await;
1175 assert!(!radio.link_up("A"), "{:#?}", radio.events());
1176 })
1177 .await;
1178 }
1179
1180 #[tokio::test(start_paused = true)]
1181 async fn disconnect_error_still_marks_state_disconnected() {
1182 within(LIMIT, async {
1183 let radio = FakeRadio::new();
1184 let core = connected_core(&radio, ReconnectOptions::default()).await;
1185 radio.fail_disconnects("A");
1186
1187 let result = core.disconnect().await;
1188
1189 assert!(matches!(result, Err(Error::Timeout { .. })), "{result:?}");
1190 assert_eq!(core.state().await, ConnectionState::Disconnected);
1191 assert!(core.link().await.is_none());
1192 assert!(!radio.link_up("A"));
1193 })
1194 .await;
1195 }
1196
1197 #[tokio::test(start_paused = true)]
1201 async fn gives_up_after_max_attempts() {
1202 within(LIMIT, async {
1203 for (max_delay, total) in [
1204 (Duration::from_secs(60), Duration::from_secs(120)),
1205 (Duration::MAX, Duration::MAX),
1206 ] {
1207 let radio = FakeRadio::new();
1208 let options = ReconnectOptions::default()
1209 .max_attempts(2)
1210 .max_delay(max_delay);
1211 let core = connected_core(&radio, options).await;
1212 radio.script_connects("A", [false; 3]);
1213 radio.lose_link("A");
1214
1215 let result = core.run(run_op).await;
1216
1217 assert!(
1218 matches!(&result, Err(Error::Timeout { operation, duration }) if operation.contains("reconnect to 'A'") && *duration == total),
1219 "{result:?}"
1220 );
1221 assert_eq!(core.state().await, ConnectionState::Failed);
1222 assert_eq!(radio.connect_count("A"), 3, "the first link and two attempts");
1223 radio.assert_no_drop_teardown();
1224 }
1225 })
1226 .await;
1227 }
1228
1229 #[tokio::test(start_paused = true)]
1230 async fn queued_callers_share_a_failed_recovery() {
1231 within(LIMIT, async {
1232 let radio = FakeRadio::new();
1233 let options = ReconnectOptions::default().max_attempts(2);
1234 let core = Arc::new(connected_core(&radio, options).await);
1235 radio.script_connects("A", [false; 4]);
1236 radio.lose_link("A");
1237 let start = Instant::now();
1238
1239 let recovering = spawn_run(&core);
1240 sleep(Duration::from_millis(500)).await;
1241 let waiting = spawn_run(&core);
1243 let reconnecting = tokio::spawn({
1244 let core = Arc::clone(&core);
1245 async move { core.reconnect().await }
1246 });
1247
1248 let recovering = recovering.await.expect("join");
1249 let waiting = waiting.await.expect("join");
1250 let reconnecting = reconnecting.await.expect("join");
1251 assert!(
1252 matches!(&recovering, Err(Error::Timeout { operation, .. }) if operation.contains("reconnect to 'A'")),
1253 "{recovering:?}"
1254 );
1255 assert!(matches!(waiting, Err(Error::NotConnected)), "{waiting:?}");
1256 assert!(
1257 matches!(reconnecting, Err(Error::NotConnected)),
1258 "{reconnecting:?}"
1259 );
1260 assert_eq!(
1261 radio.connect_count("A"),
1262 3,
1263 "the first link and one recovery's two attempts: {:#?}",
1264 radio.events()
1265 );
1266 assert_eq!(start.elapsed(), Duration::from_secs(3), "1 s + 2 s of backoff");
1267
1268 let later = core.run(run_op).await;
1270 assert!(matches!(later, Err(Error::Timeout { .. })), "{later:?}");
1271 assert_eq!(radio.connect_count("A"), 5);
1272 assert_eq!(core.state().await, ConnectionState::Failed);
1273 })
1274 .await;
1275 }
1276
1277 #[tokio::test(start_paused = true)]
1278 async fn disconnect_while_closing_the_old_link_starts_no_attempt() {
1279 within(LIMIT, async {
1280 let radio = FakeRadio::new();
1281 let core = Arc::new(connected_core(&radio, ReconnectOptions::default()).await);
1282 radio.lose_link("A");
1283
1284 let recovering = spawn_run(&core);
1289 let disconnecting = tokio::spawn({
1290 let core = Arc::clone(&core);
1291 async move { core.disconnect().await }
1292 });
1293
1294 let recovering = recovering.await.expect("join");
1295 let disconnected = disconnecting.await.expect("join");
1296 assert!(
1297 matches!(recovering, Err(Error::Cancelled)),
1298 "{recovering:?}"
1299 );
1300 assert!(disconnected.is_ok(), "{disconnected:?}");
1301 assert_eq!(core.attempt_count(), 0, "{:#?}", radio.events());
1302 assert_eq!(core.state().await, ConnectionState::Disconnected);
1303 assert_eq!(radio.connect_count("A"), 1);
1304 assert!(!radio.link_up("A"));
1305 })
1306 .await;
1307 }
1308
1309 fn link_events(radio: &FakeRadio) -> Vec<(Duration, FakeEvent)> {
1311 radio
1312 .events()
1313 .into_iter()
1314 .filter(|(_, event)| {
1315 matches!(
1316 event,
1317 FakeEvent::Connected { .. } | FakeEvent::Disconnect { .. }
1318 )
1319 })
1320 .collect()
1321 }
1322
1323 fn connected_at(secs: u64, handle: u64) -> (Duration, FakeEvent) {
1324 let id = "A".into();
1325 (
1326 Duration::from_secs(secs),
1327 FakeEvent::Connected { id, handle },
1328 )
1329 }
1330
1331 fn disconnected_at(secs: u64, handle: u64) -> (Duration, FakeEvent) {
1332 let id = "A".into();
1333 (
1334 Duration::from_secs(secs),
1335 FakeEvent::Disconnect { id, handle },
1336 )
1337 }
1338
1339 #[tokio::test(start_paused = true)]
1346 async fn queued_callers_recover_after_a_recovery_dropped_while_closing_the_old_link() {
1347 within(LIMIT, async {
1348 let radio = FakeRadio::new();
1349 let core = Arc::new(connected_core(&radio, ReconnectOptions::default()).await);
1350 radio.set_disconnect_delay("A", Duration::from_secs(4));
1351
1352 let in_flight = tokio::spawn({
1355 let core = Arc::clone(&core);
1356 async move {
1357 core.run(|link| {
1358 Box::pin(async move {
1359 sleep(Duration::from_secs(1)).await;
1360 link.op().await
1361 })
1362 })
1363 .await
1364 }
1365 });
1366 sleep(Duration::from_millis(500)).await;
1367 radio.lose_link("A");
1368
1369 let recovering = tokio::spawn({
1373 let core = Arc::clone(&core);
1374 async move {
1375 tokio::time::timeout(Duration::from_millis(1500), core.run(run_op)).await
1376 }
1377 });
1378 sleep(Duration::from_secs(1)).await;
1379 let arriving = spawn_run(&core);
1380
1381 let recovering = recovering.await.expect("join");
1382 assert!(recovering.is_err(), "{recovering:?}");
1383 let in_flight = in_flight.await.expect("join");
1384 let arriving = arriving.await.expect("join");
1385 assert!(in_flight.is_ok(), "{in_flight:?} after {:#?}", radio.events());
1386 assert!(arriving.is_ok(), "{arriving:?} after {:#?}", radio.events());
1387 assert_eq!(
1390 link_events(&radio),
1391 [connected_at(0, 1), disconnected_at(5, 1), connected_at(6, 2)]
1392 );
1393 assert_eq!(radio.connect_count("A"), 2, "one shared reconnect");
1394 assert_eq!(core.link().await.map(|link| link.handle()), Some(2));
1395 assert_eq!(core.state().await, ConnectionState::Connected);
1396
1397 sleep(Duration::from_secs(60)).await;
1399 assert!(radio.link_up("A"), "{:#?}", radio.events());
1400 radio.assert_no_drop_teardown();
1401 })
1402 .await;
1403 }
1404
1405 #[tokio::test(start_paused = true)]
1409 async fn a_recovery_waits_for_the_close_of_a_dropped_disconnect() {
1410 within(LIMIT, async {
1411 let radio = FakeRadio::new();
1412 let core = connected_core(&radio, ReconnectOptions::default()).await;
1413 radio.set_disconnect_delay("A", Duration::from_secs(4));
1414
1415 let disconnected_in_time =
1416 tokio::time::timeout(Duration::from_secs(1), core.disconnect()).await;
1417 assert!(disconnected_in_time.is_err(), "{disconnected_in_time:?}");
1418 let result = core.run(run_op).await;
1420
1421 assert!(result.is_ok(), "{result:?} after {:#?}", radio.events());
1422 assert_eq!(
1423 link_events(&radio),
1424 [
1425 connected_at(0, 1),
1426 disconnected_at(4, 1),
1427 connected_at(5, 2)
1428 ]
1429 );
1430 sleep(Duration::from_secs(60)).await;
1431 assert!(radio.link_up("A"), "{:#?}", radio.events());
1432 radio.assert_no_drop_teardown();
1433 })
1434 .await;
1435 }
1436
1437 #[tokio::test(start_paused = true)]
1442 async fn a_recovery_waits_for_the_new_link_that_a_dropped_recovery_closes() {
1443 within(LIMIT, async {
1444 let radio = FakeRadio::new();
1445 let core = Arc::new(connected_core(&radio, ReconnectOptions::default()).await);
1446 radio.set_disconnect_delay("A", Duration::from_secs(4));
1447 radio.set_connect_delay("A", Duration::from_secs(10));
1448 radio.lose_link("A");
1449
1450 let recovering = tokio::spawn({
1454 let core = Arc::clone(&core);
1455 async move { tokio::time::timeout(Duration::from_secs(17), core.run(run_op)).await }
1456 });
1457 sleep(Duration::from_secs(10)).await;
1458 let reading = core.link.read().await;
1459 sleep(Duration::from_secs(6)).await;
1460 core.cancel();
1461 radio.set_connect_delay("A", Duration::ZERO);
1462 drop(reading);
1464
1465 let recovering = recovering.await.expect("join");
1466 assert!(recovering.is_err(), "{recovering:?}");
1467 core.reset_cancellation();
1468 let result = core.run(run_op).await;
1469
1470 assert!(result.is_ok(), "{result:?} after {:#?}", radio.events());
1471 assert_eq!(
1472 link_events(&radio),
1473 [
1474 connected_at(0, 1),
1475 disconnected_at(4, 1),
1476 connected_at(15, 2),
1477 disconnected_at(20, 2),
1478 connected_at(21, 3),
1479 ]
1480 );
1481 sleep(Duration::from_secs(60)).await;
1482 assert!(radio.link_up("A"), "{:#?}", radio.events());
1483 radio.assert_no_drop_teardown();
1484 })
1485 .await;
1486 }
1487
1488 #[tokio::test(start_paused = true)]
1493 async fn run_owned_runs_again_only_after_a_connection_error_and_at_most_twice() {
1494 within(LIMIT, async {
1495 let radio = FakeRadio::new();
1496 let core = connected_core(&radio, ReconnectOptions::default()).await;
1497 let calls = AtomicU32::new(0);
1498
1499 let handle = core
1500 .run_owned(|link| {
1501 let first = calls.fetch_add(1, Ordering::SeqCst) == 0;
1502 let handle = link.handle();
1503 async move {
1504 if first {
1505 Err(Error::NotConnected)
1506 } else {
1507 Ok(handle)
1508 }
1509 }
1510 })
1511 .await;
1512 assert_eq!(handle.ok(), Some(2), "{:#?}", radio.events());
1513 assert_eq!(calls.load(Ordering::SeqCst), 2);
1514 assert_eq!(radio.connect_count("A"), 2);
1515
1516 calls.store(0, Ordering::SeqCst);
1517 let invalid = core
1518 .run_owned(|_| {
1519 calls.fetch_add(1, Ordering::SeqCst);
1520 async { Err::<(), _>(Error::InvalidData("bad".into())) }
1521 })
1522 .await;
1523 assert!(
1524 matches!(&invalid, Err(Error::InvalidData(msg)) if msg == "bad"),
1525 "{invalid:?}"
1526 );
1527 assert_eq!(calls.load(Ordering::SeqCst), 1, "{:#?}", radio.events());
1528 assert_eq!(radio.connect_count("A"), 2);
1529
1530 calls.store(0, Ordering::SeqCst);
1531 let lost = core
1532 .run_owned(|_| {
1533 calls.fetch_add(1, Ordering::SeqCst);
1534 async { Err::<(), _>(Error::NotConnected) }
1535 })
1536 .await;
1537 assert!(matches!(lost, Err(Error::NotConnected)), "{lost:?}");
1538 assert_eq!(calls.load(Ordering::SeqCst), 2, "{:#?}", radio.events());
1539 assert_eq!(radio.connect_count("A"), 3, "one reconnect");
1540 radio.assert_no_drop_teardown();
1541 })
1542 .await;
1543 }
1544
1545 #[tokio::test(start_paused = true)]
1553 async fn a_stop_after_the_connect_returns_closes_the_new_link() {
1554 within(LIMIT, async {
1555 for stop in ["cancel_reconnect", "disconnect"] {
1556 let radio = FakeRadio::new();
1557 let core = Arc::new(connected_core(&radio, ReconnectOptions::default()).await);
1558 radio.set_connect_delay("A", Duration::from_secs(10));
1559 radio.lose_link("A");
1560
1561 let recovering = spawn_run(&core);
1565 sleep(Duration::from_secs(5)).await;
1566 let reading = core.link.read().await;
1567 sleep(Duration::from_secs(10)).await;
1568 assert_eq!(radio.connect_count("A"), 2, "{stop}: {:#?}", radio.events());
1569 let disconnecting = if stop == "disconnect" {
1570 let core = Arc::clone(&core);
1571 Some(tokio::spawn(async move { core.disconnect().await }))
1572 } else {
1573 core.cancel();
1574 None
1575 };
1576 sleep(Duration::from_millis(10)).await;
1578 drop(reading);
1579
1580 let result = recovering.await.expect("join");
1581 assert!(
1582 matches!(result, Err(Error::Cancelled)),
1583 "{stop}: {result:?}"
1584 );
1585 if let Some(disconnecting) = disconnecting {
1586 let disconnected = disconnecting.await.expect("join");
1587 assert!(disconnected.is_ok(), "{disconnected:?}");
1588 }
1589 let events = radio.events();
1590 assert!(
1591 core.link().await.is_none(),
1592 "{stop}: the new link was installed: {events:#?}"
1593 );
1594 assert_eq!(core.state().await, ConnectionState::Disconnected, "{stop}");
1595 assert!(
1596 position(
1597 &events,
1598 &FakeEvent::Disconnect {
1599 id: "A".into(),
1600 handle: 2,
1601 }
1602 )
1603 .is_some(),
1604 "{stop}: the new link was not closed: {events:#?}"
1605 );
1606 assert!(!radio.link_up("A"), "{stop}: {events:#?}");
1607 radio.assert_no_drop_teardown();
1608 }
1609 })
1610 .await;
1611 }
1612}